From ca58a61d67170064dfb6d448e462f7aaee850e01 Mon Sep 17 00:00:00 2001 From: Yaniv Israel Date: Sun, 28 Jun 2026 15:21:34 +0300 Subject: [PATCH] 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). --- Makefile | 3 +- litellm/integrations/custom_guardrail.py | 4 +- .../guardrail_hooks/deepkeep/__init__.py | 4 +- .../guardrail_hooks/deepkeep/deepkeep.py | 49 +++++------------ .../proxy/guardrails/guardrail_registry.py | 8 +-- .../guardrails/guardrail_hooks/deepkeep.py | 4 +- scripts/type_check_gate.py | 53 +++++++++++++------ 7 files changed, 55 insertions(+), 70 deletions(-) diff --git a/Makefile b/Makefile index 2f8d7f2d62d..164e292ad46 100644 --- a/Makefile +++ b/Makefile @@ -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 diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index f08ab3536af..353ece8da29 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -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 diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/__init__.py index 57937ed00aa..2fb113de1eb 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/__init__.py @@ -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, diff --git a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py index 76c3f701049..b7783db1775 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py +++ b/litellm/proxy/guardrails/guardrail_hooks/deepkeep/deepkeep.py @@ -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: diff --git a/litellm/proxy/guardrails/guardrail_registry.py b/litellm/proxy/guardrails/guardrail_registry.py index ef60a1b7110..dc4f0649331 100644 --- a/litellm/proxy/guardrails/guardrail_registry.py +++ b/litellm/proxy/guardrails/guardrail_registry.py @@ -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. diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/deepkeep.py b/litellm/types/proxy/guardrails/guardrail_hooks/deepkeep.py index 76e055ee7d1..fcbb779ddf4 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/deepkeep.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/deepkeep.py @@ -16,9 +16,7 @@ class DeepKeepGuardrailConfigModelOptionalParams(BaseModel): ) -class DeepKeepGuardrailConfigModel( - GuardrailConfigModel[DeepKeepGuardrailConfigModelOptionalParams] -): +class DeepKeepGuardrailConfigModel(GuardrailConfigModel[DeepKeepGuardrailConfigModelOptionalParams]): api_key: Optional[str] = Field( default=None, description=( diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py index 2ef332d91ea..5571480e4a0 100644 --- a/scripts/type_check_gate.py +++ b/scripts/type_check_gate.py @@ -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 = "" @@ -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__":