fix: merge upstream/litellm_internal_staging (197 commits), resolve conflicts

7 conflicts resolved:
- 6 Python files: upstream added new code with old-style typing (Optional,
  Dict, List) on lines where we had ruff-fixed modern syntax (str | None,
  dict, list). Took upstream's version then re-ran ruff UP006/UP045/F401
  --fix to keep both the new content and ruff compliance.
- test_openapi_compliance.py: upstream replaced 'role' with 'steps' in
  output_fields and updated the spec comment. Took upstream's version.

Also: added _resolve_base() fallback to type_check_gate.py and removed
the hard 'git fetch origin litellm_internal_staging' from the Makefile's
lint-basedpyright target (same pattern as ruff_strict_gate.py fix).
This commit is contained in:
Yaniv Israel 2026-06-28 15:21:34 +03:00
parent c11c99d51a
commit ca58a61d67
7 changed files with 55 additions and 70 deletions

View file

@ -127,8 +127,7 @@ lint-ruff-FULL-dev: install-dev
else echo "No changed .py files to check."; fi
lint-basedpyright: install-dev
git fetch origin litellm_internal_staging
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --base origin/litellm_internal_staging
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py
lint-basedpyright-budget-update: install-dev
($(UV_RUN) basedpyright --outputjson || true) | $(UV_RUN) python scripts/type_check_gate.py --update

View file

@ -474,9 +474,7 @@ class CustomGuardrail(CustomLogger):
return True
return False
async def async_pre_call_deployment_hook(
self, kwargs: dict[str, Any], call_type: CallTypes | None
) -> dict | None:
async def async_pre_call_deployment_hook(self, kwargs: dict[str, Any], call_type: CallTypes | None) -> dict | None:
from litellm.proxy._types import UserAPIKeyAuth
# should run guardrail

View file

@ -15,9 +15,7 @@ def initialize_guardrail(litellm_params: "LitellmParams", guardrail: "Guardrail"
api_base=litellm_params.api_base,
api_key=litellm_params.api_key,
firewall_id=getattr(litellm_params, "deepkeep_firewall_id", None),
unreachable_fallback=getattr(
litellm_params, "unreachable_fallback", "fail_closed"
),
unreachable_fallback=getattr(litellm_params, "unreachable_fallback", "fail_closed"),
extra_headers=getattr(litellm_params, "extra_headers", None),
guardrail_name=guardrail.get("guardrail_name", ""),
event_hook=litellm_params.mode,

View file

@ -75,9 +75,7 @@ class DeepKeepGuardrail(CustomGuardrail):
extra_headers: dict[str, str] | None = None,
**kwargs: Any,
):
self.async_handler = get_async_httpx_client(
llm_provider=httpxSpecialProvider.GuardrailCallback
)
self.async_handler = get_async_httpx_client(llm_provider=httpxSpecialProvider.GuardrailCallback)
# API key
deepkeep_api_key = api_key or os.environ.get("DEEPKEEP_API_KEY")
@ -111,9 +109,7 @@ class DeepKeepGuardrail(CustomGuardrail):
else:
self.api_base = f"{base_url}{_DEEPKEEP_GUARDRAIL_ENDPOINT}"
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = (
unreachable_fallback
)
self.unreachable_fallback: Literal["fail_closed", "fail_open"] = unreachable_fallback
self.extra_headers: dict[str, str] = extra_headers or {}
# Set supported event hooks
@ -169,10 +165,7 @@ class DeepKeepGuardrail(CustomGuardrail):
result_metadata[key] = value
# Handle the token → hash alias (only when no explicit hash was provided)
if (
metadata_dict.get("user_api_key_token") is not None
and "user_api_key_hash" not in result_metadata
):
if metadata_dict.get("user_api_key_token") is not None and "user_api_key_hash" not in result_metadata:
result_metadata["user_api_key_hash"] = metadata_dict["user_api_key_token"]
return result_metadata
@ -197,9 +190,7 @@ class DeepKeepGuardrail(CustomGuardrail):
http_status_code: int | None = None,
) -> GenericGuardrailAPIInputs:
"""Allow the request to proceed when the guardrail is unreachable (fail-open mode)."""
status_suffix = (
f" http_status_code={http_status_code}" if http_status_code else ""
)
status_suffix = f" http_status_code={http_status_code}" if http_status_code else ""
verbose_proxy_logger.critical(
"DeepKeep guardrail unreachable (fail-open). Proceeding without guardrail.%s "
"guardrail_name=%s api_base=%s input_type=%s litellm_call_id=%s litellm_trace_id=%s",
@ -225,9 +216,7 @@ class DeepKeepGuardrail(CustomGuardrail):
) -> GenericGuardrailAPIInputs:
"""Handle errors from the DeepKeep API with fail-open/fail-closed logic."""
if is_unreachable and self.unreachable_fallback == "fail_open":
http_status_code = getattr(
getattr(error, "response", None), "status_code", None
)
http_status_code = getattr(getattr(error, "response", None), "status_code", None)
return self._fail_open_passthrough(
inputs=inputs,
input_type=input_type,
@ -300,9 +289,7 @@ class DeepKeepGuardrail(CustomGuardrail):
GuardrailRaisedException: If the guardrail blocks the request.
DeepKeepGuardrailAPIError: If the API call fails (in fail-closed mode).
"""
verbose_proxy_logger.debug(
"DeepKeep guardrail: applying guardrail, input_type=%s", input_type
)
verbose_proxy_logger.debug("DeepKeep guardrail: applying guardrail, input_type=%s", input_type)
texts = inputs.get("texts", [])
images = inputs.get("images")
@ -320,9 +307,7 @@ class DeepKeepGuardrail(CustomGuardrail):
additional_params: dict[str, Any] = {"firewall_id": self.firewall_id}
dynamic_params = self.get_guardrail_dynamic_request_body_params(request_body)
if dynamic_params:
additional_params.update(
{k: v for k, v in dynamic_params.items() if k != "firewall_id"}
)
additional_params.update({k: v for k, v in dynamic_params.items() if k != "firewall_id"})
# Extract user API key metadata
user_metadata = self._extract_user_api_key_metadata(request_data)
@ -360,12 +345,8 @@ class DeepKeepGuardrail(CustomGuardrail):
action = response_json.get("action", "NONE")
if action == "BLOCKED":
error_message = (
response_json.get("blocked_reason") or "Content violates policy"
)
verbose_proxy_logger.warning(
"DeepKeep guardrail blocked request: %s", error_message
)
error_message = response_json.get("blocked_reason") or "Content violates policy"
verbose_proxy_logger.warning("DeepKeep guardrail blocked request: %s", error_message)
raise GuardrailRaisedException(
guardrail_name=GUARDRAIL_NAME,
message=error_message,
@ -384,9 +365,7 @@ class DeepKeepGuardrail(CustomGuardrail):
except GuardrailRaisedException:
raise
except Timeout as e:
return self._handle_guardrail_request_error(
e, inputs, input_type, logging_obj
)
return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)
except httpx.HTTPStatusError as e:
status_code = getattr(getattr(e, "response", None), "status_code", None)
is_unreachable = status_code in (502, 503, 504)
@ -394,13 +373,9 @@ class DeepKeepGuardrail(CustomGuardrail):
e, inputs, input_type, logging_obj, is_unreachable=is_unreachable
)
except httpx.RequestError as e:
return self._handle_guardrail_request_error(
e, inputs, input_type, logging_obj
)
return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj)
except Exception as e:
return self._handle_guardrail_request_error(
e, inputs, input_type, logging_obj, is_unreachable=False
)
return self._handle_guardrail_request_error(e, inputs, input_type, logging_obj, is_unreachable=False)
@staticmethod
def get_config_model() -> type | None:

View file

@ -349,9 +349,7 @@ class GuardrailRegistry:
except Exception as e:
raise Exception(f"Error getting guardrail from DB: {str(e)}")
async def get_guardrail_by_name_from_db(
self, guardrail_name: str, prisma_client: PrismaClient
) -> Guardrail | None:
async def get_guardrail_by_name_from_db(self, guardrail_name: str, prisma_client: PrismaClient) -> Guardrail | None:
"""
Get a guardrail by its name from the database
"""
@ -703,9 +701,7 @@ class InMemoryGuardrailHandler:
# Initialize fresh (will add new callback to litellm.callbacks)
return self.initialize_guardrail(guardrail=guardrail, config_file_path=config_file_path, source=source)
def sync_guardrail_from_db(
self, guardrail: Guardrail, config_file_path: str | None = None
) -> Guardrail | None:
def sync_guardrail_from_db(self, guardrail: Guardrail, config_file_path: str | None = None) -> Guardrail | None:
"""
Sync a guardrail from DB - initializes if new, re-initializes if changed.
This is the method to call during DB polling.

View file

@ -16,9 +16,7 @@ class DeepKeepGuardrailConfigModelOptionalParams(BaseModel):
)
class DeepKeepGuardrailConfigModel(
GuardrailConfigModel[DeepKeepGuardrailConfigModelOptionalParams]
):
class DeepKeepGuardrailConfigModel(GuardrailConfigModel[DeepKeepGuardrailConfigModelOptionalParams]):
api_key: Optional[str] = Field(
default=None,
description=(

View file

@ -40,6 +40,7 @@ REPO_ROOT = Path(__file__).resolve().parent.parent
BUDGET_PATH = REPO_ROOT / "basedpyright-code-budget.json"
PYRIGHT_CONFIG = REPO_ROOT / "pyrightconfig.json"
DEFAULT_BASE = "origin/litellm_internal_staging"
_FALLBACK_BASE = "upstream/litellm_internal_staging"
# Bucket for a basedpyright diagnostic with no `rule`. Counted so it's gated.
UNCODED = "<uncoded>"
@ -128,9 +129,7 @@ def base_counts(ref: str) -> dict[str, int]:
exe = shutil.which("basedpyright") or "basedpyright"
with _temp_worktree(ref) as worktree:
shutil.copy(PYRIGHT_CONFIG, worktree / "pyrightconfig.json")
proc = subprocess.run(
[exe, "--outputjson"], cwd=worktree, capture_output=True, text=True
)
proc = subprocess.run([exe, "--outputjson"], cwd=worktree, capture_output=True, text=True)
return count_basedpyright(proc.stdout, root=worktree)
@ -149,9 +148,7 @@ def evaluate(
return sorted(breaches)
def is_vacuous_run(
counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]]
) -> bool:
def is_vacuous_run(counts: Mapping[str, int], budget: Mapping[str, Mapping[str, int]]) -> bool:
"""True when nothing was parsed but the budget expects errors -- the
signature of a type checker that crashed or produced no output. The CI pipe
swallows the tool's exit code (`tool || true`), so without this guard an
@ -164,16 +161,12 @@ def cmd_update(counts: Mapping[str, int]) -> None:
budget = {
code: {
"baseline": count,
"slack": (
existing[code]["slack"] if code in existing else _seed_slack(count)
),
"slack": (existing[code]["slack"] if code in existing else _seed_slack(count)),
}
for code, count in sorted(counts.items())
}
BUDGET_PATH.write_text(json.dumps(budget, indent=2, sort_keys=True) + "\n")
print(
f"Re-captured basedpyright per-rule budget: {len(budget)} rules, {sum(counts.values())} errors total"
)
print(f"Re-captured basedpyright per-rule budget: {len(budget)} rules, {sum(counts.values())} errors total")
def cmd_check(base_ref: str) -> None:
@ -204,9 +197,7 @@ def cmd_check(base_ref: str) -> None:
return
print("FAIL: basedpyright errors exceed the per-rule ceiling:")
for breach in breaches:
print(
f" {breach.code}: total {breach.total} over cap {breach.cap} (this change added {breach.added})"
)
print(f" {breach.code}: total {breach.total} over cap {breach.cap} (this change added {breach.added})")
print(
"Reduce the new errors or remove an equal number elsewhere; the ceiling is "
"baseline + slack in basedpyright-code-budget.json."
@ -216,6 +207,36 @@ def cmd_check(base_ref: str) -> None:
raise SystemExit(1)
def _resolve_base(ref: str) -> str:
"""Return ref if it resolves locally; fall back to the upstream/ equivalent.
Forks with a different 'origin' (e.g. Azure DevOps) won't have
'origin/litellm_internal_staging', so we try 'upstream/...' before
giving up.
"""
result = subprocess.run(
["git", "rev-parse", "--verify", ref],
cwd=REPO_ROOT,
capture_output=True,
)
if result.returncode == 0:
return ref
fallback = ref.replace("origin/", "upstream/", 1)
if fallback != ref:
fb = subprocess.run(
["git", "rev-parse", "--verify", fallback],
cwd=REPO_ROOT,
capture_output=True,
)
if fb.returncode == 0:
print(
f"Note: '{ref}' not found, using '{fallback}' as base.",
file=sys.stderr,
)
return fallback
return ref # let the caller fail with a clear error
def main() -> None:
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("--base", default=DEFAULT_BASE)
@ -224,7 +245,7 @@ def main() -> None:
if args.update:
cmd_update(count_basedpyright(sys.stdin.read()))
else:
cmd_check(args.base)
cmd_check(_resolve_base(args.base))
if __name__ == "__main__":