mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
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:
parent
c11c99d51a
commit
ca58a61d67
7 changed files with 55 additions and 70 deletions
3
Makefile
3
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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -16,9 +16,7 @@ class DeepKeepGuardrailConfigModelOptionalParams(BaseModel):
|
|||
)
|
||||
|
||||
|
||||
class DeepKeepGuardrailConfigModel(
|
||||
GuardrailConfigModel[DeepKeepGuardrailConfigModelOptionalParams]
|
||||
):
|
||||
class DeepKeepGuardrailConfigModel(GuardrailConfigModel[DeepKeepGuardrailConfigModelOptionalParams]):
|
||||
api_key: Optional[str] = Field(
|
||||
default=None,
|
||||
description=(
|
||||
|
|
|
|||
|
|
@ -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__":
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue