From 1a7023366fa112670201d3ae13e8ad6ad7418ad8 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Sat, 3 Oct 2026 16:07:34 -0700 Subject: [PATCH 01/18] feat(roi): measure shipping velocity, quality, and recorded spend (#44426) * feat(ui): prototype observed engineering ROI dashboard * feat(roi): replace effort estimates with measured repository metrics * fix(roi): finish connection recovery and generated API contracts * fix(roi): show merged changes before accounts are linked * fix(roi): preserve selected report tab across refreshes * fix(roi): recover app authorization and keep detail values readable * fix(roi): reuse the shared OAuth HTTP client * feat(roi): combine providers and compare equal reporting periods * docs: explain ROI metrics for first-time readers * fix(roi): preserve connections and scheduled reports during setup * ci(roi): assign database contracts to the active Postgres shard * fix(roi): preserve issue counts and normalized connections * fix(ui): compact ROI dashboard header and metrics * fix(ui): show ROI repository count with expandable list * fix(ui): wrap ROI controls within narrow panels * fix(roi): restore sample report preview and simplify setup --- .github/scripts/assert_ci_coverage.py | 65 +- .github/workflows/test-postgres.yml | 20 +- litellm/proxy/_lazy_openapi_snapshot.json | 1416 +++++++++++++++++ .../roi_calculator_endpoints.py | 243 ++- .../roi_observed_endpoints.py | 587 +++++++ litellm/proxy/roi_calculator/README.md | 52 + litellm/proxy/roi_calculator/branch_spend.py | 21 +- litellm/proxy/roi_calculator/github.py | 74 +- .../proxy/roi_calculator/github_observed.py | 196 +++ litellm/proxy/roi_calculator/gitlab.py | 47 +- litellm/proxy/roi_calculator/oauth.py | 401 +++++ .../roi_calculator/observed_analytics.py | 191 +++ litellm/proxy/roi_calculator/observed_sync.py | 265 +++ .../roi_calculator/observed_workspace.py | 135 ++ litellm/proxy/roi_calculator/settings.py | 253 +++ litellm/proxy/roi_calculator/sync_store.py | 37 +- litellm/repositories/config_repository.py | 24 + litellm/types/roi_calculator.py | 15 +- litellm/types/roi_observed.py | 213 +++ tests/integration/README.md | 4 +- tests/integration/conftest.py | 6 +- .../integration/database/test_roi_observed.py | 823 ++++++++++ tests/integration/run.py | 2 + .../spend/test_roi_branch_spend.py | 90 +- .../test_roi_calculator_endpoints.py | 22 +- .../roi_calculator/test_github_observed.py | 148 ++ .../unit/proxy/roi_calculator/test_gitlab.py | 25 + tests/unit/proxy/roi_calculator/test_oauth.py | 25 + .../roi_calculator/test_observed_analytics.py | 176 ++ .../roi_calculator/test_observed_sync.py | 140 ++ .../roi_calculator/test_observed_workspace.py | 136 ++ tests/unit/test_assert_ci_coverage.py | 13 + .../ObservedAccounts.integration.test.tsx | 88 + .../_components/ObservedAccounts.tsx | 209 +++ .../ObservedConnections.integration.test.tsx | 198 +++ .../_components/ObservedConnections.tsx | 606 +++++++ .../_components/ObservedDetails.tsx | 232 +++ .../ObservedROIView.integration.test.tsx | 305 ++++ .../_components/ObservedROIView.tsx | 253 +++ .../_components/ObservedReport.tsx | 611 +++++++ .../_components/observedData.test.ts | 110 ++ .../_components/observedData.ts | 212 +++ .../_components/observedDemo.test.ts | 42 + .../_components/observedDemo.ts | 120 ++ .../_components/useObservedReport.ts | 60 + .../app/(dashboard)/roi-calculator/page.tsx | 7 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 781 +++++++++ 47 files changed, 9483 insertions(+), 216 deletions(-) create mode 100644 litellm/proxy/management_endpoints/roi_observed_endpoints.py create mode 100644 litellm/proxy/roi_calculator/README.md create mode 100644 litellm/proxy/roi_calculator/github_observed.py create mode 100644 litellm/proxy/roi_calculator/oauth.py create mode 100644 litellm/proxy/roi_calculator/observed_analytics.py create mode 100644 litellm/proxy/roi_calculator/observed_sync.py create mode 100644 litellm/proxy/roi_calculator/observed_workspace.py create mode 100644 litellm/proxy/roi_calculator/settings.py create mode 100644 litellm/types/roi_observed.py create mode 100644 tests/integration/database/test_roi_observed.py create mode 100644 tests/unit/proxy/roi_calculator/test_github_observed.py create mode 100644 tests/unit/proxy/roi_calculator/test_oauth.py create mode 100644 tests/unit/proxy/roi_calculator/test_observed_analytics.py create mode 100644 tests/unit/proxy/roi_calculator/test_observed_sync.py create mode 100644 tests/unit/proxy/roi_calculator/test_observed_workspace.py create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.integration.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.integration.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedReport.tsx create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedData.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedData.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedDemo.test.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedDemo.ts create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/useObservedReport.ts diff --git a/.github/scripts/assert_ci_coverage.py b/.github/scripts/assert_ci_coverage.py index 3022f94a599..24cb314cd27 100644 --- a/.github/scripts/assert_ci_coverage.py +++ b/.github/scripts/assert_ci_coverage.py @@ -130,37 +130,24 @@ def _unit_selection_arms(repo_root: pathlib.Path = REPO_ROOT) -> Mapping[str, fr text: Final = _uncommented(script.read_text()) return MappingProxyType( { - label: frozenset( - match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(body) - ) + label: frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(body)) for label, body in SELECTION_ARM_RE.findall(text) } ) def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]: - return frozenset( - token for tokens in _unit_selection_arms(repo_root).values() for token in tokens - ) + return frozenset(token for tokens in _unit_selection_arms(repo_root).values() for token in tokens) def _wired_unit_flags(scalars: Iterable[Scalar]) -> frozenset[str]: - return frozenset( - scalar.value - for scalar in scalars - if scalar.key == "unit-flag" and "${{" not in scalar.value - ) + return frozenset(scalar.value for scalar in scalars if scalar.key == "unit-flag" and "${{" not in scalar.value) -def _shard_tokens( - scalars: Iterable[Scalar], arms: Mapping[str, frozenset[str]] -) -> frozenset[str]: +def _shard_tokens(scalars: Iterable[Scalar], arms: Mapping[str, frozenset[str]]) -> frozenset[str]: wired: Final = _wired_unit_flags(scalars) return _invoked_test_tokens(scalars) | frozenset( - token - for label, tokens in arms.items() - if label in wired - for token in tokens + token for label, tokens in arms.items() if label in wired for token in tokens ) @@ -544,17 +531,37 @@ def _integration_groups(runner: pathlib.Path) -> dict[str, tuple[str, ...]]: return {group: tuple(folders) for group, folders in ast.literal_eval(mapping).items()} +def _integration_github_files(runner: pathlib.Path) -> frozenset[str]: + module: Final = ast.parse(runner.read_text()) + literal: Final = next( + ( + node.value + for node in module.body + if isinstance(node, ast.AnnAssign) + and isinstance(node.target, ast.Name) + and node.target.id == "GITHUB_FILES" + ), + None, + ) + if literal is None: + return frozenset() + values: Final = literal.args[0] if isinstance(literal, ast.Call) else literal + return frozenset(ast.literal_eval(values)) + + def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozenset[str], tuple[Finding, ...]]: runner: Final = repo_root / "tests/integration/run.py" if not runner.exists(): return frozenset(), () groups: Final = _integration_groups(runner) + github_files: Final = _integration_github_files(runner) integration_root: Final = repo_root / "tests/integration" paths: Final = frozenset( str(path.relative_to(repo_root)) for folders in groups.values() for folder in folders for path in (integration_root / folder).rglob("test_*.py") + if str(path.relative_to(repo_root)) not in github_files ) browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json" browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else () @@ -595,10 +602,22 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens for path in (repo_root / ".github/workflows").glob("*.y*ml") for scalar in _scalars(yaml.safe_load(path.read_text()), path.name) ) - findings: Final = tuple( - Finding(path, "integration contract is also selected by GitHub Actions") - for path in paths - if any(_token_covers(token, path) for token in gha_tokens) + findings: Final = ( + tuple( + Finding(path, "integration contract is also selected by GitHub Actions") + for path in paths + if any(_token_covers(token, path) for token in gha_tokens) + ) + + tuple( + Finding(path, "GitHub-owned integration contract has no invoking workflow") + for path in sorted(github_files) + if not any(_token_covers(token, path) for token in gha_tokens) + ) + + tuple( + Finding(path, "GitHub-owned integration file is missing") + for path in sorted(github_files) + if not (repo_root / path).is_file() + ) ) browser_commands: Final = tuple( scalar.value @@ -642,7 +661,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens return frozenset(), findings + ( Finding(str(runner.relative_to(repo_root)), "dedicated CircleCI runner is missing"), ) - return paths | browser_paths, findings + group_findings + browser_findings + exclusion_findings + return paths | browser_paths | github_files, findings + group_findings + browser_findings + exclusion_findings def main() -> int: diff --git a/.github/workflows/test-postgres.yml b/.github/workflows/test-postgres.yml index 519d387976e..a1c639acbeb 100644 --- a/.github/workflows/test-postgres.yml +++ b/.github/workflows/test-postgres.yml @@ -45,6 +45,13 @@ jobs: fail-fast: false matrix: include: + - shard: roi-database + test-path: "tests/integration/database/test_roi_observed.py" + seed: none + workers: 0 + timeout-minutes: 10 + job-timeout-minutes: 35 + - shard: proxy-behavior test-path: "tests/proxy_behavior" seed: db-push @@ -147,7 +154,7 @@ jobs: env: TEST_PATH: ${{ matrix.test-path }} WORKERS: ${{ matrix.workers }} - PYTEST_ADDOPTS: ${{ matrix.shard == 'proxy-behavior' && '--cov=./litellm --cov-report=xml:coverage-lens-postgres.xml' || '' }} + PYTEST_ADDOPTS: ${{ matrix.shard == 'proxy-behavior' && '--cov=./litellm --cov-report=xml:coverage-lens-postgres.xml' || matrix.shard == 'roi-database' && '--cov=./litellm --cov-report=xml:coverage-roi-postgres.xml' || '' }} run: | if [ "${WORKERS}" = "0" ]; then uv run --no-sync pytest ${TEST_PATH:?} -vv --tb=short --durations=10 @@ -165,3 +172,14 @@ jobs: files: coverage-lens-postgres.xml flags: lens-postgres fail_ci_if_error: true + + - name: Upload ROI database coverage + if: steps.changes.outputs.decision != 'skip' && matrix.shard == 'roi-database' && !cancelled() + uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1 + with: + use_oidc: true + version: v11.3.1 + root_dir: ${{ github.workspace }} + files: coverage-roi-postgres.xml + flags: roi-postgres + fail_ci_if_error: true diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 4cf6db6f257..7b9af2f2b98 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -49014,6 +49014,963 @@ "title": "HTTPValidationError", "type": "object" }, + "ObservedAccount": { + "properties": { + "connection_id": { + "title": "Connection Id", + "type": "string" + }, + "login": { + "title": "Login", + "type": "string" + } + }, + "required": [ + "connection_id", + "login" + ], + "title": "ObservedAccount", + "type": "object" + }, + "ObservedApp": { + "properties": { + "api_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Api Url" + }, + "callback_url": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Callback Url" + }, + "can_install": { + "default": false, + "title": "Can Install", + "type": "boolean" + }, + "configured": { + "title": "Configured", + "type": "boolean" + } + }, + "required": [ + "configured" + ], + "title": "ObservedApp", + "type": "object" + }, + "ObservedApps": { + "properties": { + "github": { + "$ref": "#/components/schemas/ObservedApp" + }, + "gitlab": { + "$ref": "#/components/schemas/ObservedApp" + } + }, + "required": [ + "github", + "gitlab" + ], + "title": "ObservedApps", + "type": "object" + }, + "ObservedAuthorization": { + "properties": { + "url": { + "title": "Url", + "type": "string" + } + }, + "required": [ + "url" + ], + "title": "ObservedAuthorization", + "type": "object" + }, + "ObservedConnection": { + "properties": { + "api_url": { + "title": "Api Url", + "type": "string" + }, + "connection_type": { + "enum": [ + "token", + "app" + ], + "title": "Connection Type", + "type": "string" + }, + "has_token": { + "title": "Has Token", + "type": "boolean" + }, + "id": { + "default": "", + "title": "Id", + "type": "string" + }, + "ready": { + "title": "Ready", + "type": "boolean" + }, + "repos": { + "items": { + "type": "string" + }, + "title": "Repos", + "type": "array" + }, + "source_provider": { + "enum": [ + "github", + "gitlab" + ], + "title": "Source Provider", + "type": "string" + }, + "update_interval_minutes": { + "title": "Update Interval Minutes", + "type": "number" + } + }, + "required": [ + "source_provider", + "api_url", + "repos", + "has_token", + "update_interval_minutes", + "ready", + "connection_type" + ], + "title": "ObservedConnection", + "type": "object" + }, + "ObservedConnectionIdentities": { + "properties": { + "api_url": { + "title": "Api Url", + "type": "string" + }, + "id": { + "title": "Id", + "type": "string" + }, + "identity_map": { + "additionalProperties": { + "type": "string" + }, + "title": "Identity Map", + "type": "object" + }, + "repos": { + "items": { + "type": "string" + }, + "title": "Repos", + "type": "array" + }, + "source_provider": { + "enum": [ + "github", + "gitlab" + ], + "title": "Source Provider", + "type": "string" + }, + "unmatched_logins": { + "items": { + "type": "string" + }, + "title": "Unmatched Logins", + "type": "array" + } + }, + "required": [ + "id", + "source_provider", + "api_url", + "repos", + "identity_map", + "unmatched_logins" + ], + "title": "ObservedConnectionIdentities", + "type": "object" + }, + "ObservedHumanSummary": { + "properties": { + "median_merge_hours": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Median Merge Hours" + } + }, + "required": [ + "median_merge_hours" + ], + "title": "ObservedHumanSummary", + "type": "object" + }, + "ObservedIdentities": { + "properties": { + "connections": { + "default": [], + "items": { + "$ref": "#/components/schemas/ObservedConnectionIdentities" + }, + "title": "Connections", + "type": "array" + }, + "gateway_emails": { + "items": { + "type": "string" + }, + "title": "Gateway Emails", + "type": "array" + }, + "identity_map": { + "additionalProperties": { + "type": "string" + }, + "title": "Identity Map", + "type": "object" + }, + "unmatched_logins": { + "items": { + "type": "string" + }, + "title": "Unmatched Logins", + "type": "array" + } + }, + "required": [ + "gateway_emails", + "identity_map", + "unmatched_logins" + ], + "title": "ObservedIdentities", + "type": "object" + }, + "ObservedIdentityUpdate": { + "additionalProperties": false, + "properties": { + "accounts": { + "anyOf": [ + { + "items": { + "$ref": "#/components/schemas/ObservedAccount" + }, + "maxItems": 500, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Accounts" + }, + "email": { + "title": "Email", + "type": "string" + }, + "logins": { + "default": [], + "items": { + "type": "string" + }, + "maxItems": 100, + "title": "Logins", + "type": "array" + } + }, + "required": [ + "email" + ], + "title": "ObservedIdentityUpdate", + "type": "object" + }, + "ObservedPeriod": { + "properties": { + "agent_authored": { + "title": "Agent Authored", + "type": "integer" + }, + "agents_without_requester": { + "title": "Agents Without Requester", + "type": "integer" + }, + "explicitly_titled_revert_prs": { + "title": "Explicitly Titled Revert Prs", + "type": "integer" + }, + "human_authored": { + "title": "Human Authored", + "type": "integer" + }, + "human_summary": { + "$ref": "#/components/schemas/ObservedHumanSummary" + }, + "matched_internal_prs": { + "title": "Matched Internal Prs", + "type": "integer" + }, + "matched_users_recorded_spend": { + "title": "Matched Users Recorded Spend", + "type": "number" + }, + "median_merge_hours": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Median Merge Hours" + }, + "merged_prs": { + "title": "Merged Prs", + "type": "integer" + }, + "missing_author": { + "title": "Missing Author", + "type": "integer" + }, + "new_bug_labeled_issues": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "New Bug Labeled Issues" + }, + "new_regression_labeled_issues": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "New Regression Labeled Issues" + }, + "spend_observation": { + "enum": [ + "records_present", + "no_records" + ], + "title": "Spend Observation", + "type": "string" + }, + "window": { + "$ref": "#/components/schemas/ObservedWindow" + } + }, + "required": [ + "window", + "merged_prs", + "median_merge_hours", + "human_authored", + "agent_authored", + "missing_author", + "agents_without_requester", + "matched_internal_prs", + "new_bug_labeled_issues", + "new_regression_labeled_issues", + "explicitly_titled_revert_prs", + "matched_users_recorded_spend", + "spend_observation", + "human_summary" + ], + "title": "ObservedPeriod", + "type": "object" + }, + "ObservedPeriods": { + "properties": { + "current": { + "$ref": "#/components/schemas/ObservedPeriod" + }, + "last_year": { + "$ref": "#/components/schemas/ObservedPeriod" + }, + "previous": { + "$ref": "#/components/schemas/ObservedPeriod" + } + }, + "required": [ + "current", + "previous", + "last_year" + ], + "title": "ObservedPeriods", + "type": "object" + }, + "ObservedPerson": { + "properties": { + "accounts": { + "default": [], + "items": { + "$ref": "#/components/schemas/ObservedAccount" + }, + "title": "Accounts", + "type": "array" + }, + "email": { + "title": "Email", + "type": "string" + }, + "logins": { + "items": { + "type": "string" + }, + "title": "Logins", + "type": "array" + }, + "name": { + "title": "Name", + "type": "string" + }, + "periods": { + "$ref": "#/components/schemas/ObservedPersonPeriods" + } + }, + "required": [ + "name", + "email", + "logins", + "periods" + ], + "title": "ObservedPerson", + "type": "object" + }, + "ObservedPersonPeriod": { + "properties": { + "declared_agent_owned": { + "title": "Declared Agent Owned", + "type": "integer" + }, + "direct_authored": { + "title": "Direct Authored", + "type": "integer" + }, + "gateway_recorded_spend": { + "title": "Gateway Recorded Spend", + "type": "number" + }, + "median_merge_hours": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Median Merge Hours" + }, + "merged_prs": { + "title": "Merged Prs", + "type": "integer" + }, + "pr_urls": { + "items": { + "type": "string" + }, + "title": "Pr Urls", + "type": "array" + }, + "prs_per_week": { + "title": "Prs Per Week", + "type": "number" + }, + "recorded_spend_per_attributed_pr": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Recorded Spend Per Attributed Pr" + }, + "spend_observation": { + "enum": [ + "records_present", + "no_records" + ], + "title": "Spend Observation", + "type": "string" + } + }, + "required": [ + "merged_prs", + "prs_per_week", + "median_merge_hours", + "direct_authored", + "declared_agent_owned", + "gateway_recorded_spend", + "recorded_spend_per_attributed_pr", + "spend_observation", + "pr_urls" + ], + "title": "ObservedPersonPeriod", + "type": "object" + }, + "ObservedPersonPeriods": { + "properties": { + "current": { + "$ref": "#/components/schemas/ObservedPersonPeriod" + }, + "last_year": { + "$ref": "#/components/schemas/ObservedPersonPeriod" + }, + "previous": { + "$ref": "#/components/schemas/ObservedPersonPeriod" + } + }, + "required": [ + "current", + "previous", + "last_year" + ], + "title": "ObservedPersonPeriods", + "type": "object" + }, + "ObservedPullPeriods": { + "properties": { + "current": { + "items": { + "$ref": "#/components/schemas/ObservedPullResponse" + }, + "title": "Current", + "type": "array" + }, + "last_year": { + "items": { + "$ref": "#/components/schemas/ObservedPullResponse" + }, + "title": "Last Year", + "type": "array" + }, + "previous": { + "items": { + "$ref": "#/components/schemas/ObservedPullResponse" + }, + "title": "Previous", + "type": "array" + } + }, + "required": [ + "current", + "previous", + "last_year" + ], + "title": "ObservedPullPeriods", + "type": "object" + }, + "ObservedPullResponse": { + "properties": { + "agent": { + "default": false, + "title": "Agent", + "type": "boolean" + }, + "author": { + "title": "Author", + "type": "string" + }, + "branch_cost": { + "$ref": "#/components/schemas/ROIBranchAttribution" + }, + "connection_id": { + "default": "", + "title": "Connection Id", + "type": "string" + }, + "created_at": { + "anyOf": [ + { + "format": "date-time", + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Created At" + }, + "merge_hours": { + "anyOf": [ + { + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Merge Hours" + }, + "merged_at": { + "format": "date-time", + "title": "Merged At", + "type": "string" + }, + "number": { + "title": "Number", + "type": "integer" + }, + "profile_email": { + "default": "", + "title": "Profile Email", + "type": "string" + }, + "repo": { + "title": "Repo", + "type": "string" + }, + "requester": { + "default": "", + "title": "Requester", + "type": "string" + }, + "source_branch": { + "default": "", + "title": "Source Branch", + "type": "string" + }, + "source_repo": { + "default": "", + "title": "Source Repo", + "type": "string" + }, + "title": { + "title": "Title", + "type": "string" + }, + "url": { + "title": "Url", + "type": "string" + } + }, + "required": [ + "repo", + "number", + "title", + "url", + "author", + "merged_at", + "merge_hours", + "branch_cost" + ], + "title": "ObservedPullResponse", + "type": "object" + }, + "ObservedReport": { + "properties": { + "captured_at": { + "format": "date-time", + "title": "Captured At", + "type": "string" + }, + "connections": { + "default": [], + "items": { + "$ref": "#/components/schemas/ObservedSource" + }, + "title": "Connections", + "type": "array" + }, + "people": { + "items": { + "$ref": "#/components/schemas/ObservedPerson" + }, + "title": "People", + "type": "array" + }, + "periods": { + "$ref": "#/components/schemas/ObservedPeriods" + }, + "pulls": { + "$ref": "#/components/schemas/ObservedPullPeriods" + }, + "repos": { + "items": { + "type": "string" + }, + "title": "Repos", + "type": "array" + }, + "source_provider": { + "enum": [ + "github", + "gitlab", + "mixed" + ], + "title": "Source Provider", + "type": "string" + }, + "unlinked_branches": { + "items": { + "$ref": "#/components/schemas/ROIBranchSpend" + }, + "title": "Unlinked Branches", + "type": "array" + }, + "unmatched_logins": { + "items": { + "type": "string" + }, + "title": "Unmatched Logins", + "type": "array" + } + }, + "required": [ + "source_provider", + "repos", + "captured_at", + "periods", + "people", + "pulls", + "unlinked_branches", + "unmatched_logins" + ], + "title": "ObservedReport", + "type": "object" + }, + "ObservedReportResponse": { + "properties": { + "report": { + "anyOf": [ + { + "$ref": "#/components/schemas/ObservedReport" + }, + { + "type": "null" + } + ] + } + }, + "required": [ + "report" + ], + "title": "ObservedReportResponse", + "type": "object" + }, + "ObservedSettings": { + "properties": { + "api_url": { + "title": "Api Url", + "type": "string" + }, + "connection_type": { + "enum": [ + "token", + "app" + ], + "title": "Connection Type", + "type": "string" + }, + "connections": { + "default": [], + "items": { + "$ref": "#/components/schemas/ObservedConnection" + }, + "title": "Connections", + "type": "array" + }, + "has_token": { + "title": "Has Token", + "type": "boolean" + }, + "id": { + "default": "", + "title": "Id", + "type": "string" + }, + "ready": { + "title": "Ready", + "type": "boolean" + }, + "repos": { + "items": { + "type": "string" + }, + "title": "Repos", + "type": "array" + }, + "source_provider": { + "enum": [ + "github", + "gitlab" + ], + "title": "Source Provider", + "type": "string" + }, + "update_interval_minutes": { + "title": "Update Interval Minutes", + "type": "number" + } + }, + "required": [ + "source_provider", + "api_url", + "repos", + "has_token", + "update_interval_minutes", + "ready", + "connection_type" + ], + "title": "ObservedSettings", + "type": "object" + }, + "ObservedSettingsUpdate": { + "additionalProperties": false, + "properties": { + "api_url": { + "title": "Api Url", + "type": "string" + }, + "connection_id": { + "anyOf": [ + { + "maxLength": 100, + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Connection Id" + }, + "repos": { + "items": { + "type": "string" + }, + "title": "Repos", + "type": "array" + }, + "source_provider": { + "enum": [ + "github", + "gitlab" + ], + "title": "Source Provider", + "type": "string" + }, + "token": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Token" + }, + "update_interval_minutes": { + "anyOf": [ + { + "maximum": 43200.0, + "minimum": 0.0, + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Update Interval Minutes" + } + }, + "required": [ + "source_provider", + "api_url", + "repos" + ], + "title": "ObservedSettingsUpdate", + "type": "object" + }, + "ObservedSource": { + "properties": { + "api_url": { + "title": "Api Url", + "type": "string" + }, + "id": { + "title": "Id", + "type": "string" + }, + "repos": { + "items": { + "type": "string" + }, + "title": "Repos", + "type": "array" + }, + "source_provider": { + "enum": [ + "github", + "gitlab" + ], + "title": "Source Provider", + "type": "string" + } + }, + "required": [ + "id", + "source_provider", + "api_url", + "repos" + ], + "title": "ObservedSource", + "type": "object" + }, + "ObservedWindow": { + "properties": { + "end": { + "format": "date", + "title": "End", + "type": "string" + }, + "start": { + "format": "date", + "title": "Start", + "type": "string" + } + }, + "required": [ + "start", + "end" + ], + "title": "ObservedWindow", + "type": "object" + }, "ROIBranchAttribution": { "properties": { "branch": { @@ -49703,6 +50660,15 @@ "title": "Ready", "type": "boolean" }, + "report_mode": { + "default": "legacy", + "enum": [ + "legacy", + "observed" + ], + "title": "Report Mode", + "type": "string" + }, "repos": { "items": { "type": "string" @@ -49834,6 +50800,21 @@ ], "title": "Gitlab Token" }, + "report_mode": { + "anyOf": [ + { + "enum": [ + "legacy", + "observed" + ], + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Report Mode" + }, "repos": { "anyOf": [ { @@ -50267,6 +51248,441 @@ ] } }, + "/roi-calculator/observed/apps": { + "get": { + "operationId": "observed_apps_roi_calculator_observed_apps_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedApps" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Observed Apps", + "tags": [ + "roi_calculator" + ] + } + }, + "/roi-calculator/observed/identities": { + "get": { + "operationId": "get_observed_identities_roi_calculator_observed_identities_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedIdentities" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Observed Identities", + "tags": [ + "roi_calculator" + ] + }, + "put": { + "operationId": "save_observed_identities_roi_calculator_observed_identities_put", + "requestBody": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedIdentityUpdate" + } + } + }, + "required": true + }, + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedReportResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Save Observed Identities", + "tags": [ + "roi_calculator" + ] + } + }, + "/roi-calculator/observed/oauth/{provider}/start": { + "post": { + "operationId": "start_observed_authorization_roi_calculator_observed_oauth__provider__start_post", + "parameters": [ + { + "in": "path", + "name": "provider", + "required": true, + "schema": { + "enum": [ + "github", + "gitlab" + ], + "title": "Provider", + "type": "string" + } + }, + { + "in": "query", + "name": "install", + "required": false, + "schema": { + "default": false, + "title": "Install", + "type": "boolean" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedAuthorization" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Start Observed Authorization", + "tags": [ + "roi_calculator" + ] + } + }, + "/roi-calculator/observed/report": { + "get": { + "operationId": "get_observed_report_roi_calculator_observed_report_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedReportResponse" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Observed Report", + "tags": [ + "roi_calculator" + ] + } + }, + "/roi-calculator/observed/repositories": { + "get": { + "operationId": "observed_repositories_roi_calculator_observed_repositories_get", + "parameters": [ + { + "in": "query", + "name": "connection", + "required": false, + "schema": { + "anyOf": [ + { + "maxLength": 100, + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Connection" + } + }, + { + "in": "query", + "name": "query", + "required": false, + "schema": { + "default": "", + "maxLength": 200, + "title": "Query", + "type": "string" + } + }, + { + "in": "query", + "name": "page", + "required": false, + "schema": { + "default": 1, + "maximum": 1000, + "minimum": 1, + "title": "Page", + "type": "integer" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ROIRepositoriesResponse" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Observed Repositories", + "tags": [ + "roi_calculator" + ] + } + }, + "/roi-calculator/observed/settings": { + "get": { + "operationId": "get_observed_settings_roi_calculator_observed_settings_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedSettings" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Observed Settings", + "tags": [ + "roi_calculator" + ] + }, + "put": { + "operationId": "save_observed_settings_roi_calculator_observed_settings_put", + "requestBody": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedSettingsUpdate" + } + } + }, + "required": true + }, + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ObservedSettings" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Save Observed Settings", + "tags": [ + "roi_calculator" + ] + } + }, + "/roi-calculator/observed/sync": { + "delete": { + "operationId": "cancel_observed_sync_roi_calculator_observed_sync_delete", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ROISyncStatus" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cancel Observed Sync", + "tags": [ + "roi_calculator" + ] + }, + "get": { + "operationId": "get_observed_sync_roi_calculator_observed_sync_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ROISyncStatus" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Observed Sync", + "tags": [ + "roi_calculator" + ] + }, + "post": { + "operationId": "start_observed_sync_roi_calculator_observed_sync_post", + "parameters": [ + { + "in": "query", + "name": "days", + "required": false, + "schema": { + "anyOf": [ + { + "maximum": 366, + "minimum": 1, + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Days" + } + } + ], + "responses": { + "202": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ROISyncStatus" + } + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Start Observed Sync", + "tags": [ + "roi_calculator" + ] + } + }, "/roi-calculator/report": { "get": { "operationId": "get_roi_calculator_report_roi_calculator_report_get", diff --git a/litellm/proxy/management_endpoints/roi_calculator_endpoints.py b/litellm/proxy/management_endpoints/roi_calculator_endpoints.py index bb6d60db723..4d2f53f2227 100644 --- a/litellm/proxy/management_endpoints/roi_calculator_endpoints.py +++ b/litellm/proxy/management_endpoints/roi_calculator_endpoints.py @@ -15,20 +15,29 @@ from apscheduler.schedulers.asyncio import ( # pyright: ignore[reportMissingTyp AsyncIOScheduler, ) from fastapi import APIRouter, Depends, FastAPI, HTTPException, Query -from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError +from pydantic import BaseModel, ConfigDict, SecretStr, TypeAdapter, ValidationError from starlette.types import Receive, Scope, Send from litellm.llms.custom_httpx.http_handler import ( AsyncHTTPHandler, get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # shared client factory has untyped params ) -from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth -from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper +from litellm.proxy.management_endpoints.roi_observed_endpoints import router as observed_router from litellm.proxy.roi_calculator.analytics import normalize_email, summarize from litellm.proxy.roi_calculator.branch_spend import BranchSpendDatabase, read_branch_spend from litellm.proxy.roi_calculator.estimator import CompletionCaller, EstimatorModel from litellm.proxy.roi_calculator.github import SourceError +from litellm.proxy.roi_calculator.settings import ( + active_connection, + get_roi_config_repository, + load_settings, + load_stored_settings, + read_admin, + save_settings, + write_admin, +) from litellm.proxy.roi_calculator.source import create_source from litellm.proxy.roi_calculator.sync import ( BranchSpendReader, @@ -62,29 +71,13 @@ from litellm.types.roi_calculator import ( ) router: Final = APIRouter() +router.include_router(observed_router) _SETTINGS_KEY: Final = "roi_calculator_settings" _REPORT_KEY: Final = "roi_calculator_report" _SYNC_MANAGER: Final = SyncManager() _ROI_TAGS: Final[list[str | Enum]] = ["roi calculator"] # mutable-ok: FastAPI requires list-valued route tags -class _StoredSettings(BaseModel): - model_config = ConfigDict(extra="ignore") - - source_provider: Literal["github", "gitlab"] = "github" - gitlab_api_url: str = "https://gitlab.com/api/v4" - gitlab_token: str = "" - github_api_url: str = "https://api.github.com" - github_token: str = "" - estimator_key: str = "" - repos: tuple[str, ...] = () - estimator_model: str = "" - estimator_prompt: str = DEFAULT_PROMPT - backfill_days: int = Field(default=7, ge=1, le=3650) - update_interval_minutes: float = Field(default=1440, ge=0, le=43200) - identity_map: Mapping[str, str] = Field(default_factory=lambda: MappingProxyType({})) - - class _RouterEstimatorParams(BaseModel): model_config = ConfigDict(extra="ignore", from_attributes=True) @@ -108,38 +101,6 @@ class _RouterEstimatorDeployment(BaseModel): model_info: _RouterEstimatorModelInfo | None = None -async def _read_admin( - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], -) -> UserAPIKeyAuth: - if user_api_key_dict.user_role not in ( - LitellmUserRoles.PROXY_ADMIN, - LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, - ): - raise HTTPException(status_code=403, detail="Only proxy admins can access the ROI Calculator.") - return user_api_key_dict - - -async def _write_admin( - user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], -) -> UserAPIKeyAuth: - if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: - raise HTTPException(status_code=403, detail="Only proxy admins can change ROI Calculator settings.") - return user_api_key_dict - - -async def get_roi_config_repository( - _user: Annotated[UserAPIKeyAuth, Depends(_read_admin)], -) -> ConfigRepository: - from litellm.proxy.proxy_server import prisma_client - - if prisma_client is None: - raise HTTPException( - status_code=500, - detail=CommonProxyErrors.db_not_connected_error.value, - ) - return ConfigRepository(prisma_client, use_writer=True) - - def get_roi_sync_manager() -> SyncManager: return _SYNC_MANAGER @@ -222,66 +183,6 @@ def _router_estimator_choices() -> tuple[ROIEstimatorModel, ...]: return tuple(choice for choice in choices if choice.model_name in names) -async def _load_stored_settings(repository: ConfigRepository) -> _StoredSettings: - parameter: Final = await repository.get_param(_SETTINGS_KEY) - if parameter is None: - return _StoredSettings() - try: - return _StoredSettings.model_validate(parameter.param_value) - except ValidationError: - raise HTTPException(status_code=500, detail="Stored ROI Calculator settings are invalid.") from None - - -async def _load_settings(repository: ConfigRepository) -> ROISettings: - stored: Final = await _load_stored_settings(repository) - token: Final = decrypt_value_helper(stored.github_token, _SETTINGS_KEY) if stored.github_token else "" - try: - return ROISettings( - source_provider=stored.source_provider, - gitlab_api_url=stored.gitlab_api_url, - gitlab_token=SecretStr(decrypt_value_helper(stored.gitlab_token, _SETTINGS_KEY) or "") - if stored.gitlab_token - else SecretStr(""), - github_api_url=stored.github_api_url, - github_token=SecretStr(token or ""), - estimator_key=SecretStr(decrypt_value_helper(stored.estimator_key, _SETTINGS_KEY) or "") - if stored.estimator_key - else SecretStr(""), - update_interval_minutes=stored.update_interval_minutes, - repos=stored.repos, - estimator_model=stored.estimator_model, - estimator_prompt=stored.estimator_prompt, - backfill_days=stored.backfill_days, - identity_map=stored.identity_map, - ) - except ValidationError: - raise HTTPException(status_code=500, detail="Stored ROI Calculator settings are invalid.") from None - - -async def _save_settings( - repository: ConfigRepository, - settings: ROISettings, - encrypted_token: str, - encrypted_estimator_key: str, - encrypted_gitlab_token: str = "", -) -> None: - stored: Final = _StoredSettings( - source_provider=settings.source_provider, - gitlab_api_url=settings.gitlab_api_url, - gitlab_token=encrypted_gitlab_token, - github_api_url=settings.github_api_url, - github_token=encrypted_token, - estimator_key=encrypted_estimator_key, - update_interval_minutes=settings.update_interval_minutes, - repos=settings.repos, - estimator_model=settings.estimator_model, - estimator_prompt=settings.estimator_prompt, - backfill_days=settings.backfill_days, - identity_map=settings.identity_map, - ) - await repository.set_param(_SETTINGS_KEY, stored.model_dump(mode="json")) - - async def _load_report(repository: ConfigRepository, settings: ROISettings) -> ROIReport | None: parameter: Final = await repository.get_param(_REPORT_KEY) if parameter is None or parameter.param_value is None: @@ -302,6 +203,7 @@ def _public_settings(settings: ROISettings) -> ROISettingsResponse: choices: Final = _router_estimator_choices() models: Final = tuple(choice.model_name for choice in choices) return ROISettingsResponse( + report_mode=settings.report_mode, source_provider=settings.source_provider, gitlab_api_url=settings.gitlab_api_url, has_gitlab_token=bool(settings.gitlab_token.get_secret_value()), @@ -392,14 +294,14 @@ async def _test_estimator_access(settings: ROISettings) -> None: raise HTTPException(status_code=409, detail="The estimator key could not connect to the gateway.") from None -def _gateway_user_reader(repository: ConfigRepository) -> GatewayUserReader: +def gateway_user_reader(repository: ConfigRepository) -> GatewayUserReader: async def get_emails() -> frozenset[str]: return await read_gateway_user_emails(spend_prisma_client(repository.prisma_client)) return get_emails -def _spend_reader(repository: ConfigRepository) -> SpendReader: +def spend_reader(repository: ConfigRepository) -> SpendReader: async def get_spend(start: date, end: date) -> tuple[ROISpendRecord, ...]: prisma_client: Final = spend_prisma_client(repository.prisma_client) return await read_spend(prisma_client, start, end) @@ -407,7 +309,7 @@ def _spend_reader(repository: ConfigRepository) -> SpendReader: return get_spend -def _branch_spend_reader(repository: ConfigRepository, settings: ROISettings) -> BranchSpendReader: +def branch_spend_reader(repository: ConfigRepository, settings: ROISettings) -> BranchSpendReader: async def get_spend(start: date, end: date, repos: tuple[str, ...]) -> tuple[ROIBranchSpend, ...]: return await read_branch_spend( cast( # cast-ok: PrismaWrapper delegates methods dynamically @@ -428,10 +330,10 @@ def _branch_spend_reader(repository: ConfigRepository, settings: ROISettings) -> tags=_ROI_TAGS, ) async def get_roi_calculator_settings( - _user: Annotated[UserAPIKeyAuth, Depends(_read_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], ) -> ROISettingsResponse: - return _public_settings(await _load_settings(repository)) + return _public_settings(await load_settings(repository)) @router.put( @@ -441,11 +343,11 @@ async def get_roi_calculator_settings( ) async def update_roi_calculator_settings( patch: ROISettingsUpdate, - _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], ) -> ROISettingsResponse: - stored: Final = await _load_stored_settings(repository) - current: Final = await _load_settings(repository) + stored: Final = await load_stored_settings(repository) + current: Final = await load_settings(repository, stored) if "github_api_url" in patch.model_fields_set and patch.github_api_url is None: raise HTTPException(status_code=422, detail="GitHub API URL cannot be null.") if "gitlab_api_url" in patch.model_fields_set and patch.gitlab_api_url is None: @@ -491,6 +393,17 @@ async def update_roi_calculator_settings( ) try: settings: Final = ROISettings( + report_mode=patch.report_mode or current.report_mode, + connection_type="token" + if source_changed or token_was_supplied or "gitlab_token" in patch.model_fields_set + else current.connection_type, + oauth_refresh_token=SecretStr("") + if source_changed or token_was_supplied or "gitlab_token" in patch.model_fields_set + else current.oauth_refresh_token, + oauth_expires_at=None + if source_changed or token_was_supplied or "gitlab_token" in patch.model_fields_set + else current.oauth_expires_at, + ignored_logins=() if source_changed else current.ignored_logins, source_provider=provider, gitlab_api_url=gitlab_url, gitlab_token=SecretStr(gitlab_token), @@ -510,7 +423,15 @@ async def update_roi_calculator_settings( ) except ValidationError as exc: raise HTTPException(status_code=422, detail=exc.errors(include_context=False)) from None - await _save_settings(repository, settings, encrypted_token, encrypted_estimator_key, encrypted_gitlab) + await save_settings( + repository, + settings, + encrypted_token, + encrypted_estimator_key, + encrypted_gitlab, + revision=stored.revision, + replace_connection_id=active_connection(stored).id, + ) if source_changed: await repository.set_param(_REPORT_KEY, None) return _public_settings(settings) @@ -522,13 +443,13 @@ async def update_roi_calculator_settings( tags=_ROI_TAGS, ) async def get_roi_calculator_repositories( - _user: Annotated[UserAPIKeyAuth, Depends(_read_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], transport: Annotated[httpx.AsyncBaseTransport | None, Depends(get_github_transport)], query: Annotated[str, Query(max_length=200)] = "", page: Annotated[int, Query(ge=1, le=1000)] = 1, ) -> ROIRepositoriesResponse: - github: Final = create_source(await _load_settings(repository), transport) + github: Final = create_source(await load_settings(repository), transport) try: repos, has_more = await github.repositories(query, page) except SourceError as exc: @@ -550,12 +471,12 @@ async def get_roi_calculator_repositories( tags=_ROI_TAGS, ) async def get_roi_calculator_sync_status( - _user: Annotated[UserAPIKeyAuth, Depends(_read_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], manager: Annotated[SyncManager, Depends(get_roi_sync_manager)], ) -> ROISyncStatus: status: Final = await SyncStore(repository.prisma_client).status() or manager.status - settings: Final = await _load_settings(repository) + settings: Final = await load_settings(repository) report: Final = await _load_report(repository, settings) next_update: Final = _next_update(settings, status, report) return status.model_copy(update=MappingProxyType({"next_update": next_update.isoformat() if next_update else None})) @@ -568,25 +489,25 @@ async def get_roi_calculator_sync_status( tags=_ROI_TAGS, ) async def start_roi_calculator_sync( - _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], manager: Annotated[SyncManager, Depends(get_roi_sync_manager)], transport: Annotated[httpx.AsyncBaseTransport | None, Depends(get_github_transport)], ) -> ROISyncStatus: - settings: Final = await _load_settings(repository) + settings: Final = await load_settings(repository) public: Final = _public_settings(settings) if not public.ready: raise HTTPException(status_code=409, detail="Connect a source, select repositories, and choose a router model.") if not await manager.start( settings, repository, - _spend_reader(repository), + spend_reader(repository), _completion_caller(settings), transport, _router_estimator_models(settings.estimator_model), SyncStore(repository.prisma_client), - branch_spend_reader=_branch_spend_reader(repository, settings), - gateway_user_reader=_gateway_user_reader(repository), + branch_spend_reader=branch_spend_reader(repository, settings), + gateway_user_reader=gateway_user_reader(repository), ): raise HTTPException(status_code=409, detail="A sync is already running.") return manager.status @@ -598,7 +519,7 @@ async def start_roi_calculator_sync( tags=_ROI_TAGS, ) async def cancel_roi_calculator_sync( - _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], manager: Annotated[SyncManager, Depends(get_roi_sync_manager)], ) -> ROISyncStatus: @@ -614,7 +535,7 @@ async def cancel_roi_calculator_sync( tags=_ROI_TAGS, ) async def get_roi_calculator_report( - _user: Annotated[UserAPIKeyAuth, Depends(_read_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], mode: Literal["live", "demo"] = "live", ) -> ROIReportResponse: @@ -623,7 +544,7 @@ async def get_roi_calculator_report( sample: Final = summarize(sample_report(datetime.now(timezone.utc)), MappingProxyType({})) return ROIReportResponse(report=ROISummaryResponse.model_validate(sample)) - settings: Final = await _load_settings(repository) + settings: Final = await load_settings(repository) report: Final = await _load_report(repository, settings) if report is None: return ROIReportResponse(report=None) @@ -638,12 +559,12 @@ async def get_roi_calculator_report( ) async def update_roi_calculator_identity_map( update: ROIIdentityMapUpdate, - _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], ) -> ROIIdentityMapResponse: login: Final = update.github_login.strip().casefold() - current: Final = await _load_settings(repository) - current_stored: Final = await _load_stored_settings(repository) + current_stored: Final = await load_stored_settings(repository) + current: Final = await load_settings(repository, current_stored) try: normalize_source_login(login, current.source_provider) except ValueError as exc: @@ -657,6 +578,11 @@ async def update_roi_calculator_identity_map( else MappingProxyType({**current.identity_map, login: new_email}) ) settings: Final = ROISettings( + report_mode=current.report_mode, + connection_type=current.connection_type, + oauth_refresh_token=current.oauth_refresh_token, + oauth_expires_at=current.oauth_expires_at, + ignored_logins=current.ignored_logins, source_provider=current.source_provider, gitlab_api_url=current.gitlab_api_url, gitlab_token=current.gitlab_token, @@ -670,8 +596,13 @@ async def update_roi_calculator_identity_map( backfill_days=current.backfill_days, identity_map=identity_map, ) - await _save_settings( - repository, settings, current_stored.github_token, current_stored.estimator_key, current_stored.gitlab_token + await save_settings( + repository, + settings, + current_stored.github_token, + current_stored.estimator_key, + current_stored.gitlab_token, + revision=current_stored.revision, ) report: Final = await _load_report(repository, settings) summary: Final = summarize(report, settings.identity_map) if report is not None else None @@ -683,7 +614,8 @@ async def update_roi_calculator_identity_map( def _next_update(settings: ROISettings, status: ROISyncStatus, report: ROIReport | None) -> datetime | None: if ( - not report + settings.report_mode != "legacy" + or not report or not settings.repos or not settings.estimator_model or not settings.update_interval_minutes @@ -710,12 +642,16 @@ def register_scheduled_sync(scheduler: AsyncIOScheduler) -> None: async def run_scheduled_sync() -> None: + from litellm.proxy.management_endpoints.roi_observed_endpoints import run_observed_schedule from litellm.proxy.proxy_server import prisma_client if prisma_client is None: return repository: Final = ConfigRepository(prisma_client, use_writer=True) - settings: Final = await _load_settings(repository) + settings: Final = await load_settings(repository) + if settings.report_mode == "observed": + await run_observed_schedule() + return if not settings.update_interval_minutes or not _public_settings(settings).ready: return store: Final = SyncStore(prisma_client) @@ -727,23 +663,23 @@ async def run_scheduled_sync() -> None: await _SYNC_MANAGER.start( settings, repository, - _spend_reader(repository), + spend_reader(repository), _completion_caller(settings), estimator_models=_router_estimator_models(settings.estimator_model), coordinator=store, scheduled_interval=settings.update_interval_minutes, - branch_spend_reader=_branch_spend_reader(repository, settings), - gateway_user_reader=_gateway_user_reader(repository), + branch_spend_reader=branch_spend_reader(repository, settings), + gateway_user_reader=gateway_user_reader(repository), ) @router.post("/roi-calculator/connections/test", tags=_ROI_TAGS) async def test_roi_calculator_connections( - _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], transport: Annotated[httpx.AsyncBaseTransport | None, Depends(get_github_transport)], ) -> ROISettingsResponse: - settings: Final = await _load_settings(repository) + settings: Final = await load_settings(repository) public: Final = _public_settings(settings) if not public.ready: raise HTTPException(status_code=409, detail="Choose repositories and an available estimator model first.") @@ -760,7 +696,7 @@ async def test_roi_calculator_connections( @router.post("/roi-calculator/setup/reset", tags=_ROI_TAGS) async def reset_roi_calculator_setup( - _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], ) -> ROISettingsResponse: from uuid import uuid4 @@ -781,10 +717,17 @@ async def reset_roi_calculator_setup( if not await store.acquire(owner, status): raise HTTPException(status_code=409, detail="Cancel the running analysis before restarting setup.") try: - current: Final = await _load_settings(repository) - stored: Final = await _load_stored_settings(repository) + stored: Final = await load_stored_settings(repository) + current: Final = await load_settings(repository, stored) settings: Final = current.model_copy(update=MappingProxyType({"repos": ()})) - await _save_settings(repository, settings, stored.github_token, stored.estimator_key, stored.gitlab_token) + await save_settings( + repository, + settings, + stored.github_token, + stored.estimator_key, + stored.gitlab_token, + revision=stored.revision, + ) await store.clear_report() return _public_settings(settings) finally: diff --git a/litellm/proxy/management_endpoints/roi_observed_endpoints.py b/litellm/proxy/management_endpoints/roi_observed_endpoints.py new file mode 100644 index 00000000000..424d7af4076 --- /dev/null +++ b/litellm/proxy/management_endpoints/roi_observed_endpoints.py @@ -0,0 +1,587 @@ +from collections.abc import Mapping +from datetime import datetime, timedelta, timezone +from typing import Annotated, Final + +import httpx +from fastapi import APIRouter, Depends, HTTPException, Query, Request +from fastapi.responses import JSONResponse, RedirectResponse +from pydantic import SecretStr, TypeAdapter, ValidationError + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper +from litellm.proxy.roi_calculator.analytics import normalize_email +from litellm.proxy.roi_calculator.github import SourceError +from litellm.proxy.roi_calculator.oauth import ( + OAuthConfig, + Provider, + begin_authorization, + connected_settings, + consume_state, + exchange_code, + oauth_config, + save_grant, +) +from litellm.proxy.roi_calculator.observed_sync import ObservedSyncManager, Progress +from litellm.proxy.roi_calculator.observed_workspace import ( + collect_workspace, + scoped_data, + source_details, + summarize_workspace, +) +from litellm.proxy.roi_calculator.settings import ( + StoredConnection, + active_connection, + connection_id, + enable_observed_reporting, + get_roi_config_repository, + load_settings, + load_stored_settings, + read_admin, + save_connection_identities, + save_settings, + select_connection, + stored_connections, + write_admin, +) +from litellm.proxy.roi_calculator.source import create_source +from litellm.proxy.roi_calculator.sync import read_gateway_user_emails, spend_prisma_client +from litellm.proxy.roi_calculator.sync_store import SyncStore +from litellm.repositories.config_repository import ConfigRepository +from litellm.types.roi_calculator import ( + ROIRepositoriesResponse, + ROIRepository, + ROISettings, + ROISyncStatus, + normalize_source_login, +) +from litellm.types.roi_observed import ( + ObservedAccount, + ObservedApp, + ObservedApps, + ObservedAuthorization, + ObservedConnectionIdentities, + ObservedData, + ObservedIdentities, + ObservedIdentityUpdate, + ObservedReportResponse, + ObservedSettings, + ObservedSettingsUpdate, +) + +router: Final = APIRouter(prefix="/roi-calculator/observed", tags=["roi calculator"]) +_MANAGER: Final = ObservedSyncManager() +_REPORT_KEY: Final = "roi_observed_report" + + +def get_observed_manager() -> ObservedSyncManager: + return _MANAGER + + +def get_observed_transport() -> httpx.AsyncBaseTransport | None: + return None + + +def public_settings(settings: ROISettings) -> ObservedSettings: + return ObservedSettings( + id=connection_id(settings.source_provider, settings.source_api_url), + source_provider=settings.source_provider, + api_url=settings.source_api_url, + repos=settings.repos, + has_token=bool( + ( + settings.gitlab_token if settings.source_provider == "gitlab" else settings.github_token + ).get_secret_value() + ), + update_interval_minutes=settings.update_interval_minutes, + ready=bool(settings.repos), + connection_type=settings.connection_type, + ) + + +async def workspace_settings(repository: ConfigRepository) -> ObservedSettings: + stored: Final = await load_stored_settings(repository) + entries: Final = tuple( + [ + public_settings(await load_settings(repository, select_connection(stored, entry))) + for entry in stored_connections(stored) + ] + ) + current: Final = public_settings(await load_settings(repository, stored)) + return current.model_copy(update={"connections": entries, "ready": any(entry.ready for entry in entries)}) + + +async def _data(repository: ConfigRepository) -> ObservedData | None: + saved: Final = await repository.get_param(_REPORT_KEY) + if saved is None or saved.param_value is None: + return None + try: + data: Final = ObservedData.model_validate(saved.param_value) + except ValidationError: + raise HTTPException(500, "The saved report is invalid. Sync again to rebuild it.") from None + if data.connections: + return data + settings: Final = ROISettings.model_validate( + { + "source_provider": data.source_provider, + "repos": data.repos, + ("gitlab_api_url" if data.source_provider == "gitlab" else "github_api_url"): data.source_api_url, + } + ) + return scoped_data(data, source_details(settings)) + + +@router.get("/settings", response_model=ObservedSettings) +async def get_observed_settings( + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], +) -> ObservedSettings: + return await workspace_settings(repository) + + +@router.put("/settings", response_model=ObservedSettings) +async def save_observed_settings( + patch: ObservedSettingsUpdate, + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], + transport: Annotated[httpx.AsyncBaseTransport | None, Depends(get_observed_transport)], +) -> ObservedSettings: + if (status := await SyncStore(repository.prisma_client, "roi_observed").status()) and status.running: + raise HTTPException(409, "Cancel the running sync before changing the connection.") + original: Final = await load_stored_settings(repository) + target_id: Final = connection_id(patch.source_provider, patch.api_url) + selected_id: Final = patch.connection_id or target_id + selected: Final = next((entry for entry in stored_connections(original) if entry.id == selected_id), None) + if patch.connection_id and selected is None: + raise HTTPException(404, "This connection no longer exists. Reload Connections.") + if ( + patch.connection_id + and selected_id != target_id + and any(entry.id == target_id for entry in stored_connections(original)) + ): + raise HTTPException(409, "This provider and host are already connected. Edit that connection instead.") + if patch.token is None and selected: + await connected_settings(repository, transport, selected_id) + refreshed: Final = await load_stored_settings(repository) + saved_connection: Final = next((entry for entry in stored_connections(refreshed) if entry.id == selected_id), None) + stored: Final = select_connection(refreshed, saved_connection) if saved_connection else refreshed + current: Final = await load_settings(repository, stored) + changed: Final = target_id != connection_id(current.source_provider, current.source_api_url) + existing_token: Final = current.gitlab_token if patch.source_provider == "gitlab" else current.github_token + token: Final = patch.token if patch.token is not None else "" if changed else existing_token.get_secret_value() + updates: Final[Mapping[str, object]] = { + "report_mode": "observed", + "source_provider": patch.source_provider, + "repos": patch.repos, + "update_interval_minutes": patch.update_interval_minutes + if patch.update_interval_minutes is not None + else current.update_interval_minutes, + "identity_map": {} if changed else current.identity_map, + "ignored_logins": () if changed else current.ignored_logins, + "connection_type": "token" if patch.token is not None or changed else current.connection_type, + "oauth_refresh_token": SecretStr("") if patch.token is not None or changed else current.oauth_refresh_token, + "oauth_expires_at": None if patch.token is not None or changed else current.oauth_expires_at, + ("gitlab_api_url" if patch.source_provider == "gitlab" else "github_api_url"): patch.api_url, + ("gitlab_token" if patch.source_provider == "gitlab" else "github_token"): SecretStr(token), + } + try: + settings: Final = ROISettings.model_validate({**current.model_dump(), **updates}) + except ValidationError as exc: + raise HTTPException(422, exc.errors(include_context=False, include_input=False)) from None + source: Final = create_source(settings, transport) + try: + if settings.repos: + await source.test_repositories(settings.repos) + elif token: + await source.repositories(page=1) + except SourceError as exc: + raise HTTPException(502, str(exc)) from None + finally: + await source.close() + encrypted: Final = TypeAdapter(str).validate_python(encrypt_value_helper(token)) if token else "" + await save_settings( + repository, + settings, + encrypted if settings.source_provider == "github" else stored.github_token, + stored.estimator_key, + encrypted if settings.source_provider == "gitlab" else stored.gitlab_token, + revision=stored.revision, + replace_connection_id=patch.connection_id, + ) + return await workspace_settings(repository) + + +@router.get("/report", response_model=ObservedReportResponse) +async def get_observed_report( + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], +) -> ObservedReportResponse: + stored: Final = await load_stored_settings(repository) + data: Final = await _data(repository) + return ObservedReportResponse(report=summarize_workspace(data, stored_connections(stored)) if data else None) + + +@router.get("/identities", response_model=ObservedIdentities) +async def get_observed_identities( + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], +) -> ObservedIdentities: + stored: Final = await load_stored_settings(repository) + data: Final = await _data(repository) + report: Final = summarize_workspace(data, stored_connections(stored)) if data else None + return ObservedIdentities( + gateway_emails=tuple(sorted(await read_gateway_user_emails(spend_prisma_client(repository.prisma_client)))), + identity_map=stored.identity_map, + unmatched_logins=report.unmatched_logins if report else (), + connections=tuple( + ObservedConnectionIdentities( + id=entry.id, + source_provider=entry.source_provider, + api_url=entry.api_url, + repos=entry.repos, + identity_map=entry.identity_map, + unmatched_logins=tuple( + login.split(":", 1)[-1] for login in report.unmatched_logins if login.startswith(entry.id + ":") + ) + if report + else (), + ) + for entry in stored_connections(stored) + ), + ) + + +@router.put("/identities", response_model=ObservedReportResponse) +async def save_observed_identities( + patch: ObservedIdentityUpdate, + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], +) -> ObservedReportResponse: + email: Final = normalize_email(patch.email) + emails: Final = await read_gateway_user_emails(spend_prisma_client(repository.prisma_client)) + if email not in emails: + raise HTTPException(422, "Choose an existing internal user email.") + stored: Final = await load_stored_settings(repository) + connections: Final = stored_connections(stored) + if not connections: + raise HTTPException(409, "Connect a repository before linking accounts.") + accounts: Final = ( + patch.accounts + if patch.accounts is not None + else tuple(ObservedAccount(connection_id=active_connection(stored).id, login=login) for login in patch.logins) + ) + if any(account.connection_id not in {entry.id for entry in connections} for account in accounts): + raise HTTPException(422, "Choose an existing connection.") + data: Final = await _data(repository) + report: Final = summarize_workspace(data, connections) if data else None + existing: Final = next((person.accounts for person in report.people if person.email == email), ()) if report else () + + def update(entry: StoredConnection) -> StoredConnection: + if patch.accounts is None and entry.id != active_connection(stored).id: + return entry + try: + logins: Final = tuple( + normalize_source_login(account.login, entry.source_provider) + for account in accounts + if account.connection_id == entry.id + ) + except ValueError as exc: + raise HTTPException(422, str(exc)) from None + if any(login in entry.identity_map and entry.identity_map[login] != email for login in logins): + raise HTTPException(409, "An account is already linked to another email. Unlink it first.") + old: Final = frozenset( + ( + *(login for login, address in entry.identity_map.items() if address == email), + *(account.login for account in existing if account.connection_id == entry.id), + ) + ) + return entry.model_copy( + update={ + "identity_map": { + **{login: address for login, address in entry.identity_map.items() if address != email}, + **dict.fromkeys(logins, email), + }, + "ignored_logins": tuple(sorted((frozenset(entry.ignored_logins) | old) - frozenset(logins))), + } + ) + + updated: Final = tuple(update(entry) for entry in connections) + await save_connection_identities(repository, stored, updated) + return ObservedReportResponse(report=summarize_workspace(data, updated) if data else None) + + +@router.get("/sync", response_model=ROISyncStatus) +async def get_observed_sync( + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], + manager: Annotated[ObservedSyncManager, Depends(get_observed_manager)], +) -> ROISyncStatus: + return await SyncStore(repository.prisma_client, "roi_observed").status() or manager.status + + +async def _report_days(repository: ConfigRepository, days: int | None) -> int: + if days is not None: + return days + saved: Final = await repository.get_param(_REPORT_KEY) + if saved is None: + return 28 + try: + data: Final = ObservedData.model_validate(saved.param_value) + except ValidationError: + return 28 + return (data.current.window.end - data.current.window.start).days + 1 + + +async def _start_sync( + repository: ConfigRepository, + manager: ObservedSyncManager, + transport: httpx.AsyncBaseTransport | None, + scheduled_interval: float = 0, + days: int | None = None, +) -> bool: + from litellm.proxy.management_endpoints.roi_calculator_endpoints import branch_spend_reader + + stored: Final = await load_stored_settings(repository) + reporting_days: Final = await _report_days(repository, days) + entries: Final = tuple(entry for entry in stored_connections(stored) if entry.repos) + if not entries: + raise HTTPException(409, "Select at least one repository.") + settings: Final = tuple([await connected_settings(repository, transport, entry.id) for entry in entries]) + await enable_observed_reporting(repository) + + async def build(progress: Progress) -> ObservedData: + from litellm.proxy.management_endpoints.roi_calculator_endpoints import gateway_user_reader, spend_reader + + data: Final = await collect_workspace( + tuple((entry, branch_spend_reader(repository, entry)) for entry in settings), + spend_reader(repository), + gateway_user_reader(repository), + datetime.now(timezone.utc), + progress, + transport, + days=reporting_days, + ) + current: Final = tuple( + entry for entry in stored_connections(await load_stored_settings(repository)) if entry.repos + ) + if tuple((entry.id, entry.repos) for entry in current) != tuple((entry.id, entry.repos) for entry in entries): + raise SourceError("The connection changed during sync. Sync again with the current repositories.") + return data + + return await manager.start(build, SyncStore(repository.prisma_client, "roi_observed"), scheduled_interval) + + +@router.post("/sync", response_model=ROISyncStatus, status_code=202) +async def start_observed_sync( + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], + manager: Annotated[ObservedSyncManager, Depends(get_observed_manager)], + transport: Annotated[httpx.AsyncBaseTransport | None, Depends(get_observed_transport)], + days: Annotated[int | None, Query(ge=1, le=366)] = None, +) -> ROISyncStatus: + if not await _start_sync(repository, manager, transport, days=days): + raise HTTPException(409, "A sync is already running.") + return manager.status + + +@router.delete("/sync", response_model=ROISyncStatus) +async def cancel_observed_sync( + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], + manager: Annotated[ObservedSyncManager, Depends(get_observed_manager)], +) -> ROISyncStatus: + store: Final = SyncStore(repository.prisma_client, "roi_observed") + await store.cancel() + await manager.cancel() + return await store.status() or manager.status + + +async def run_observed_schedule() -> None: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + return + repository: Final = ConfigRepository(prisma_client, use_writer=True) + settings: Final = await load_settings(repository) + if ( + not any(entry.repos for entry in stored_connections(await load_stored_settings(repository))) + or not settings.update_interval_minutes + ): + return + status: Final = await SyncStore(prisma_client, "roi_observed").status() + if status and status.running: + return + anchor: Final = status.finished_at if status else None + if anchor and datetime.fromisoformat(anchor) + timedelta(minutes=settings.update_interval_minutes) > datetime.now( + timezone.utc + ): + return + await _start_sync(repository, _MANAGER, None, settings.update_interval_minutes) + + +@router.get("/repositories", response_model=ROIRepositoriesResponse) +async def observed_repositories( + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], + transport: Annotated[httpx.AsyncBaseTransport | None, Depends(get_observed_transport)], + connection: Annotated[str | None, Query(max_length=100)] = None, + query: Annotated[str, Query(max_length=200)] = "", + page: Annotated[int, Query(ge=1, le=1000)] = 1, +) -> ROIRepositoriesResponse: + source: Final = create_source(await connected_settings(repository, transport, connection), transport) + try: + repositories, more = await source.repositories(query, page) + except SourceError as exc: + raise HTTPException(502, str(exc)) from None + finally: + await source.close() + return ROIRepositoriesResponse( + repositories=tuple( + ROIRepository(name=name, visibility=visibility, archived=archived) + for name, visibility, archived in repositories + ), + page=page, + has_more=more, + ) + + +@router.get("/apps", response_model=ObservedApps) +async def observed_apps(_user: Annotated[UserAPIKeyAuth, Depends(read_admin)]) -> ObservedApps: + def details(provider: Provider) -> ObservedApp: + config: Final = oauth_config(provider) + return ObservedApp( + configured=config is not None, + can_install=bool(config and config.installation_url), + api_url=config.api_url if config else None, + callback_url=config.redirect_uri if config else None, + ) + + return ObservedApps(github=details("github"), gitlab=details("gitlab")) + + +@router.post("/oauth/{provider}/start", response_model=ObservedAuthorization) +async def start_observed_authorization( + provider: Provider, + _user: Annotated[UserAPIKeyAuth, Depends(write_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], + install: bool = False, +) -> JSONResponse: + config: Final = oauth_config(provider) + if config is None: + raise HTTPException(409, "Configure the provider app client ID, client secret, and PROXY_BASE_URL first.") + if install and not config.installation_url: + raise HTTPException(409, "Configure the GitHub app slug to manage repository access.") + url, nonce = await begin_authorization(repository, config, install=install) + response: Final = JSONResponse(ObservedAuthorization(url=url).model_dump(mode="json")) + response.set_cookie( + "litellm_roi_oauth", + nonce, + httponly=True, + secure=config.proxy_url.startswith("https://"), + samesite="lax", + max_age=600, + path=config.cookie_path, + ) + response.headers["Cache-Control"] = "no-store" + return response + + +def get_oauth_repository() -> ConfigRepository: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(503, "The database is unavailable. Try connecting again later.") + return ConfigRepository(prisma_client, use_writer=True) + + +def _authorization_redirect(config: OAuthConfig, query: str) -> RedirectResponse: + response: Final = RedirectResponse(config.proxy_url + "/ui/roi-calculator/?" + query, status_code=303) + response.delete_cookie("litellm_roi_oauth", path=config.cookie_path) + response.headers["Cache-Control"] = "no-store" + response.headers["Referrer-Policy"] = "no-referrer" + return response + + +@router.get("/oauth/{provider}/callback", include_in_schema=False) +async def observed_authorization_callback( + provider: Provider, + request: Request, + repository: Annotated[ConfigRepository, Depends(get_oauth_repository)], + transport: Annotated[httpx.AsyncBaseTransport | None, Depends(get_observed_transport)], + state: str = "", + code: str = "", + error: str = "", +) -> RedirectResponse: + config: Final = oauth_config(provider) + if config is None: + raise HTTPException(409, "The provider app is not configured.") + try: + return await _complete_authorization(config, request, repository, transport, state, code, error) + except HTTPException: + return _authorization_redirect(config, "connection_failed=1") + + +async def _complete_authorization( + config: OAuthConfig, + request: Request, + repository: ConfigRepository, + transport: httpx.AsyncBaseTransport | None, + state: str, + code: str, + error: str, +) -> RedirectResponse: + verified: Final = await consume_state(repository, state, request.cookies.get("litellm_roi_oauth", ""), config) + if error or not code: + return _authorization_redirect(config, "connection_cancelled=1") + if (status := await SyncStore(repository.prisma_client, "roi_observed").status()) and status.running: + raise HTTPException(409, "Cancel the running sync before changing the connection.") + grant: Final = await exchange_code(config, verified, code, transport) + validation_settings: Final = ROISettings.model_validate( + { + "source_provider": config.provider, + "connection_type": "app", + ("github_api_url" if config.provider == "github" else "gitlab_api_url"): config.api_url, + ("github_token" if config.provider == "github" else "gitlab_token"): grant.access_token, + } + ) + source: Final = create_source(validation_settings, transport) + try: + await source.repositories(page=1) + except SourceError as exc: + raise HTTPException(502, str(exc)) from None + finally: + await source.close() + await save_grant(repository, config, grant, revision=verified.settings_revision) + return _authorization_redirect(config, "connected=" + config.provider) + + +@router.get("/oauth/github/installed", include_in_schema=False) +async def observed_installation_callback( + request: Request, + repository: Annotated[ConfigRepository, Depends(get_oauth_repository)], + state: str = "", +) -> RedirectResponse: + config: Final = oauth_config("github") + if config is None or config.installation_url is None: + raise HTTPException(409, "The GitHub app is not configured.") + try: + verified: Final = await consume_state( + repository, state, request.cookies.get("litellm_roi_oauth", ""), config, flow="install" + ) + if (await load_stored_settings(repository)).revision != verified.settings_revision: + raise HTTPException(409, "The connection changed during installation. Start again from Connections.") + url, nonce = await begin_authorization(repository, config) + except HTTPException: + return _authorization_redirect(config, "connection_failed=1") + response: Final = RedirectResponse(url, status_code=303) + response.set_cookie( + "litellm_roi_oauth", + nonce, + httponly=True, + secure=config.proxy_url.startswith("https://"), + samesite="lax", + max_age=600, + path=config.cookie_path, + ) + response.headers["Cache-Control"] = "no-store" + response.headers["Referrer-Policy"] = "no-referrer" + return response diff --git a/litellm/proxy/roi_calculator/README.md b/litellm/proxy/roi_calculator/README.md new file mode 100644 index 00000000000..4db62fdad2b --- /dev/null +++ b/litellm/proxy/roi_calculator/README.md @@ -0,0 +1,52 @@ +# ROI Calculator + +The dashboard compares merged pull or merge requests, elapsed time from opening to merge, new bug and regression issues, and recorded gateway spend over 7, 28, or 90 complete UTC days. Compare against the immediately preceding period of the same length or the same-length period last year + +The calculator combines repository activity with spend recorded by the gateway. Spend per merged change is a person's recorded gateway spend during the period divided by their matched merged changes. To track a branch's AI cost, send repository and branch tags with each request + +## Connect repositories + +Use **Preview sample report** beside the title to explore the dashboard before connecting repositories. Sample periods, engineer details, quality signals, and branch spend work without changing your connections or live report. **Exit demo** returns to your report or setup + +Open `/ui/roi-calculator/`, choose GitHub or GitLab, then connect with an app or access token. Select several repositories and start the sync. Use **Add connection** to keep both providers connected. Each provider and API host retains its credentials, repositories, and identity mappings, and the report combines their activity while counting each person’s gateway spend once. Public repositories also accept an empty token, subject to the provider's anonymous API limits + +For GitHub tokens, grant read access to metadata, pull requests and issues. GitLab tokens require `read_api`. Self-hosted instances use their API URL, for example `https://git.example.com/api/v4` + +## Configure app authorization + +Register a GitHub App with read-only repository permissions for metadata, pull requests and issues. Enable expiring user access tokens and leave authorization during installation disabled, since the gateway starts authorization after installation. Generate a private key in the app settings to allow installation and store it securely. The gateway uses a generated client secret for authorization and does not need the private key + +Register a confidential GitLab OAuth application with `read_api` and `read_user` scopes + +Set `PROXY_BASE_URL` to the gateway's public URL. The callback URLs are `/roi-calculator/observed/oauth/github/callback` and `/roi-calculator/observed/oauth/gitlab/callback` + +Set the GitHub App setup URL to `/roi-calculator/observed/oauth/github/installed`, enable **Redirect on update**, and set `LITELLM_ROI_GITHUB_APP_SLUG` to its URL slug. The first connection then starts with repository installation and continues to user authorization + +Set `LITELLM_ROI_GITHUB_CLIENT_ID` and `LITELLM_ROI_GITHUB_CLIENT_SECRET` for GitHub, or `LITELLM_ROI_GITLAB_CLIENT_ID` and `LITELLM_ROI_GITLAB_CLIENT_SECRET` for GitLab. For a self-hosted provider, set `LITELLM_ROI_GITHUB_URL` or `LITELLM_ROI_GITLAB_URL` to its base URL without the API suffix + +The gateway encrypts access and refresh tokens using its configured encryption key. Authorization uses PKCE and an expiring, single-use state tied to an HTTP-only browser cookie. Refreshes are coordinated across gateway workers + +## Link people + +Use **Link accounts** to associate several current or historical usernames with one internal email. Each connection has a separate username field, so a GitHub username never matches a GitLab user implicitly. Saving immediately recalculates the report without fetching repositories again. Public profile emails match automatically when they resolve unambiguously to an internal user + +Agent-authored changes count for a person only when the supported agent metadata explicitly names a requester. Repository issue counts and revert titles are quality signals, not an individual defect score + +Bug and regression counts combine repositories with issue tracking enabled. They remain unavailable when none of the selected repositories has issue tracking enabled + +## Sync behavior + +The default refresh interval is daily and applies to every connection in the workspace. Adding or editing a connection preserves it unless `update_interval_minutes` is supplied. The observed settings API accepts `update_interval_minutes: 0` for manual updates. A cancelled or failed sync preserves the last complete report + +Existing settings retain `report_mode: legacy` and their scheduled reports until an administrator saves a connection, authorizes an app, or starts an observed sync. Reading the new dashboard alone does not change the mode. The legacy settings API can explicitly select `report_mode: legacy` again + +GitHub collection splits large searches into smaller date ranges to avoid its search-result limit. Both providers validate pagination and reject incomplete responses instead of publishing partial counts + + +## Branch request tags + +Send `repo:github.com/owner/repo` or `repo:gitlab.com/group/project` together with `branch:feature/name` in `metadata.tags`, top-level `tags`, or the comma-separated `x-litellm-tags` header. The tags must identify the source repository and branch, including forks + +The report sums recorded requests inside its UTC dates. A branch cost is assigned to a merged change only when that source branch matches one change in the period. Reused branches stay visible in Branch spend without duplicating costs across changes. No retained tagged requests means unknown cost; a recorded zero remains zero + +An empty repository produces a successful report with zero merged changes and no merge duration or spend-per-change ratio diff --git a/litellm/proxy/roi_calculator/branch_spend.py b/litellm/proxy/roi_calculator/branch_spend.py index f441683e843..d2193a6cc01 100644 --- a/litellm/proxy/roi_calculator/branch_spend.py +++ b/litellm/proxy/roi_calculator/branch_spend.py @@ -61,12 +61,23 @@ async def read_branch_spend( def attribute_branches( pulls: tuple[ROIPullRecord, ...], spend: tuple[ROIBranchSpend, ...] | None ) -> Mapping[tuple[str, int], ROIBranchAttribution]: - counts: Final = Counter((pull.get("source_repo", ""), pull.get("source_branch", "")) for pull in pulls) + return attribute_branch_keys( + tuple( + (pull["repo"], pull["number"], pull.get("source_repo", ""), pull.get("source_branch", "")) for pull in pulls + ), + spend, + ) + + +def attribute_branch_keys( + pulls: tuple[tuple[str, int, str, str], ...], spend: tuple[ROIBranchSpend, ...] | None +) -> Mapping[tuple[str, int], ROIBranchAttribution]: + counts: Final = Counter((pull[2], pull[3]) for pull in pulls) costs: Final = {(row.repo, row.branch): row for row in spend or ()} - def attribute(pull: ROIPullRecord) -> ROIBranchAttribution: - repo: Final = pull.get("source_repo", "") - branch: Final = pull.get("source_branch", "") + def attribute(pull: tuple[str, int, str, str]) -> ROIBranchAttribution: + repo: Final = pull[2] + branch: Final = pull[3] cost: Final = costs.get((repo, branch)) if spend is None: return ROIBranchAttribution(repo=repo, branch=branch, status="unavailable") @@ -78,4 +89,4 @@ def attribute_branches( repo=repo, branch=branch, spend=cost.spend, requests=cost.requests, status="matched" ) - return {(pull["repo"], pull["number"]): attribute(pull) for pull in pulls} + return {(pull[0], pull[1]): attribute(pull) for pull in pulls} diff --git a/litellm/proxy/roi_calculator/github.py b/litellm/proxy/roi_calculator/github.py index 9f8aa8c26ae..993b6b9c9fa 100644 --- a/litellm/proxy/roi_calculator/github.py +++ b/litellm/proxy/roi_calculator/github.py @@ -1,6 +1,6 @@ import asyncio from collections.abc import AsyncIterator, Mapping -from datetime import date +from datetime import date, datetime from types import MappingProxyType from typing import Final, TypeVar from urllib.parse import quote @@ -15,6 +15,7 @@ from litellm.llms.custom_httpx.http_handler import ( from litellm.proxy.roi_calculator.analytics import normalize_email from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.roi_calculator import ROIPullCommit, ROIPullEvidence, ROIPullFile, ROISettings +from litellm.types.roi_observed import ObservedIssue _T: Final = TypeVar("_T") @@ -29,6 +30,8 @@ class _GitHubModel(BaseModel): class _GitHubUser(_GitHubModel): login: str | None = None + type: str = "User" + email: str | None = None class _GitHubHeadRepository(_GitHubModel): @@ -50,6 +53,7 @@ class GitHubPullListItem(_GitHubModel): body: str | None = None head: _GitHubHead | None = None user: _GitHubUser | None = None + created_at: datetime | None = None class _RepositoryItem(_GitHubModel): @@ -215,15 +219,18 @@ _GRAPHQL_QUERY: Final = """query($owner:String!, $name:String!, $number:Int!, $c }""" -async def _request( +async def request_github( client: httpx.AsyncClient, method: str, path: str, params: Mapping[str, str | int] | None = None, json_body: object | None = None, headers: Mapping[str, str] | None = None, + *, + read_only: bool = False, ) -> httpx.Response: async def send(attempt: int) -> httpx.Response: + retryable: Final = (method == "GET" or read_only) and attempt < 2 try: response: Final = await client.request( method, @@ -233,8 +240,11 @@ async def _request( headers=headers, ) except httpx.RequestError: + if retryable: + await asyncio.sleep(0.5 * (attempt + 1)) + return await send(attempt + 1) raise SourceError("Could not reach GitHub. Check the API URL and network connection.") from None - if response.status_code in (429, 502, 503, 504) and method == "GET" and attempt < 2: + if response.status_code in (429, 502, 503, 504) and retryable: await asyncio.sleep(0.5 * (attempt + 1)) return await send(attempt + 1) if response.status_code >= 400: @@ -267,7 +277,7 @@ async def _fetch_page( headers: Mapping[str, str] | None = None, error_message: str = "GitHub returned an unexpected pagination response.", ) -> tuple[tuple[_T, ...], bool]: - response: Final = await _request( + response: Final = await request_github( client, "GET", path, @@ -312,6 +322,21 @@ class _GitHubUserProfile(_GitHubModel): email: str | None = None +class _IssueLabel(_GitHubModel): + name: str + + +class _Issue(_GitHubModel): + number: int + created_at: datetime + labels: tuple[_IssueLabel, ...] = () + pull_request: object | None = None + + +class GitHubIssueSettings(BaseModel): + has_issues: bool + + class GitHub: def __init__( self, @@ -408,8 +433,8 @@ class GitHub: async def test_repositories(self, repos: tuple[str, ...]) -> None: for repo in repos: - await _request(self.client, "GET", self._url(f"repos/{repo}"), headers=self._headers) - await _request( + await request_github(self.client, "GET", self._url(f"repos/{repo}"), headers=self._headers) + await request_github( self.client, "GET", self._url(f"repos/{repo}/pulls"), @@ -438,8 +463,40 @@ class GitHub: return await _collect(matching_pulls()) + async def issues(self, repo: str, start: date, end: date) -> tuple[ObservedIssue, ...] | None: + response: Final = await request_github(self.client, "GET", self._url(f"repos/{repo}"), headers=self._headers) + try: + settings: Final = GitHubIssueSettings.model_validate(response.json()) + except ValueError: + raise SourceError("GitHub returned invalid repository settings.") from None + if not settings.has_issues: + return None + + async def matching_issues() -> AsyncIterator[ObservedIssue]: + async for page in _pages( + self.client, + self._url(f"repos/{repo}/issues"), + TypeAdapter(tuple[_Issue, ...]), + MappingProxyType( + {"state": "all", "sort": "created", "direction": "desc", "since": f"{start}T00:00:00Z"} + ), + headers=self._headers, + ): + for issue in page: + if issue.pull_request is None and start <= issue.created_at.date() <= end: + yield ObservedIssue( + repo=repo, + number=issue.number, + created_at=issue.created_at, + labels=tuple(label.name for label in issue.labels), + ) + if page and page[-1].created_at.date() < start: + return + + return await _collect(matching_issues()) + async def evidence(self, repo: str, pull: GitHubPullListItem) -> ROIPullEvidence: - detail_response: Final = await _request( + detail_response: Final = await request_github( self.client, "GET", self._url(f"repos/{repo}/pulls/{pull.number}"), @@ -574,7 +631,7 @@ class GitHub: ) -> tuple[tuple[ROIPullCommit, ...], tuple[tuple[str, str], ...], int]: if remaining_pages == 0: raise SourceError("GitHub commit pagination limit was reached.") - response: Final = await _request( + response: Final = await request_github( self.client, "POST", endpoint, @@ -583,6 +640,7 @@ class GitHub: query=_GRAPHQL_QUERY, variables=_GraphQLVariables(owner=owner, name=name, number=number, cursor=cursor), ), + read_only=True, ) try: parsed: Final = _GRAPHQL_RESPONSE.validate_python(response.json()) diff --git a/litellm/proxy/roi_calculator/github_observed.py b/litellm/proxy/roi_calculator/github_observed.py new file mode 100644 index 00000000000..a34d34823be --- /dev/null +++ b/litellm/proxy/roi_calculator/github_observed.py @@ -0,0 +1,196 @@ +from datetime import date, datetime, time, timedelta, timezone +from typing import Final, Literal + +import httpx +from pydantic import BaseModel, Field + +from litellm.proxy.roi_calculator.github import GitHubIssueSettings, GitHubPullListItem, SourceError, request_github +from litellm.types.roi_calculator import ROISettings +from litellm.types.roi_observed import ObservedIssue + + +class _PageInfo(BaseModel): + hasNextPage: bool = False + endCursor: str | None = None + + +class _Author(BaseModel): + login: str + kind: str = Field(alias="__typename") + email: str | None = None + + +class _Repository(BaseModel): + nameWithOwner: str + + +class _Label(BaseModel): + name: str + + +class _Labels(BaseModel): + nodes: tuple[_Label, ...] = () + pageInfo: _PageInfo = Field(default_factory=_PageInfo) + + +class _Node(BaseModel): + number: int + url: str + title: str + createdAt: datetime + updatedAt: str + mergedAt: str | None = None + author: _Author | None = None + body: str = "" + headRefName: str = "" + headRefOid: str = "" + headRepository: _Repository | None = None + labels: _Labels = Field(default_factory=_Labels) + + def pull(self) -> GitHubPullListItem: + return GitHubPullListItem.model_validate( + { + "number": self.number, + "html_url": self.url, + "title": self.title, + "body": self.body, + "created_at": self.createdAt, + "merged_at": self.mergedAt, + "updated_at": self.updatedAt, + "user": {"login": self.author.login, "type": self.author.kind, "email": self.author.email} + if self.author + else None, + "head": { + "ref": self.headRefName, + "sha": self.headRefOid, + "repo": {"full_name": self.headRepository.nameWithOwner} if self.headRepository else None, + }, + } + ) + + +class _Search(BaseModel): + issueCount: int + pageInfo: _PageInfo + nodes: tuple[_Node, ...] + + +class _Data(BaseModel): + search: _Search + + +class _Response(BaseModel): + data: _Data | None = None + errors: tuple[object, ...] = () + + +_QUERY: Final = """query($q:String!, $after:String) { + search(query:$q, type:ISSUE, first:100, after:$after) { + issueCount pageInfo { hasNextPage endCursor } + nodes { + ... on PullRequest { + number url title body createdAt updatedAt mergedAt headRefName headRefOid + author { __typename login ... on User { email } } headRepository { nameWithOwner } + } + ... on Issue { + number url title createdAt updatedAt + labels(first:100) { nodes { name } pageInfo { hasNextPage } } + } + } + } +}""" + + +class GitHubObserved: + def __init__(self, settings: ROISettings, client: httpx.AsyncClient) -> None: + self._client: Final = client + self._api_url: Final = settings.github_api_url + self._url: Final = ( + "https://api.github.com/graphql" + if settings.github_api_url == "https://api.github.com" + else settings.github_api_url.removesuffix("/api/v3") + "/api/graphql" + ) + self._headers: Final = {"Authorization": "Bearer " + settings.github_token.get_secret_value()} + + async def _page(self, query: str, cursor: str | None = None) -> _Search: + response: Final = await request_github( + self._client, + "POST", + self._url, + headers=self._headers, + json_body={"query": _QUERY, "variables": {"q": query, "after": cursor}}, + read_only=True, + ) + try: + result: Final = _Response.model_validate(response.json()) + except ValueError: + raise SourceError("GitHub returned invalid activity data. Try syncing again.") from None + if result.errors or result.data is None: + raise SourceError("GitHub could not read all activity. Check app permissions and rate limits, then retry.") + return result.data.search + + async def _range( + self, repo: str, start: datetime, end: datetime, kind: Literal["pull", "issue"] + ) -> tuple[_Node, ...]: + qualifier: Final = "merged" if kind == "pull" else "created" + source: Final = "is:pr is:merged" if kind == "pull" else "is:issue" + lower: Final = start.strftime("%Y-%m-%dT%H:%M:%SZ") + upper: Final = (end - timedelta(seconds=1)).strftime("%Y-%m-%dT%H:%M:%SZ") + query: Final = f"repo:{repo} {source} {qualifier}:{lower}..{upper} sort:created-asc" + first: Final = await self._page(query) + if first.issueCount > 1000: + seconds: Final = int((end - start).total_seconds()) + if seconds < 2: + raise SourceError("GitHub has more than 1,000 results in one second. The report was not truncated.") + middle: Final = start + timedelta(seconds=seconds // 2) + left: Final = await self._range(repo, start, middle, kind) + return left + await self._range(repo, middle, end, kind) + + async def remaining(page: _Search, seen: frozenset[str]) -> tuple[_Node, ...]: + if not page.pageInfo.hasNextPage: + return page.nodes + cursor: Final = page.pageInfo.endCursor + if not cursor or cursor in seen or len(seen) >= 10: + raise SourceError("GitHub returned incomplete pagination. The previous report was kept.") + following: Final = await self._page(query, cursor) + return page.nodes + await remaining(following, seen | {cursor}) + + nodes: Final = await remaining(first, frozenset()) + if len(nodes) != first.issueCount or len(frozenset(node.url for node in nodes)) != len(nodes): + raise SourceError("GitHub activity changed during collection. Retry to get a complete report.") + return nodes + + async def _read(self, repo: str, start: date, end: date, kind: Literal["pull", "issue"]) -> tuple[_Node, ...]: + return await self._range( + repo, + datetime.combine(start, time.min, timezone.utc), + datetime.combine(end + timedelta(days=1), time.min, timezone.utc), + kind, + ) + + async def pulls(self, repo: str, start: date, end: date) -> tuple[GitHubPullListItem, ...]: + nodes: Final = await self._read(repo, start, end, "pull") + return tuple(node.pull() for node in nodes) + + async def issues(self, repo: str, start: date, end: date) -> tuple[ObservedIssue, ...] | None: + response: Final = await request_github( + self._client, "GET", f"{self._api_url}/repos/{repo}", headers=self._headers + ) + try: + settings: Final = GitHubIssueSettings.model_validate(response.json()) + except ValueError: + raise SourceError("GitHub returned invalid repository settings.") from None + if not settings.has_issues: + return None + nodes: Final = await self._read(repo, start, end, "issue") + if any(node.labels.pageInfo.hasNextPage for node in nodes): + raise SourceError("GitHub returned incomplete issue labels. The previous report was kept.") + return tuple( + ObservedIssue( + repo=repo, + number=node.number, + created_at=node.createdAt, + labels=tuple(label.name for label in node.labels.nodes), + ) + for node in nodes + ) diff --git a/litellm/proxy/roi_calculator/gitlab.py b/litellm/proxy/roi_calculator/gitlab.py index 2df35a0060b..9fbf43d8890 100644 --- a/litellm/proxy/roi_calculator/gitlab.py +++ b/litellm/proxy/roi_calculator/gitlab.py @@ -1,6 +1,6 @@ import asyncio from collections.abc import Mapping -from datetime import date +from datetime import date, datetime, timedelta from types import MappingProxyType from typing import Final, TypeVar from urllib.parse import quote @@ -16,6 +16,7 @@ from litellm.proxy.roi_calculator.github import GitHubPullListItem, SourceError from litellm.proxy.roi_calculator.source import repository_tag from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.roi_calculator import ROIPullCommit, ROIPullEvidence, ROIPullFile, ROISettings +from litellm.types.roi_observed import ObservedIssue _T: Final = TypeVar("_T", bound=BaseModel) @@ -23,6 +24,7 @@ _T: Final = TypeVar("_T", bound=BaseModel) class _User(BaseModel): username: str public_email: str | None = None + bot: bool = False class _Project(BaseModel): @@ -30,6 +32,8 @@ class _Project(BaseModel): path_with_namespace: str visibility: str = "private" archived: bool = False + issues_enabled: bool = True + issues_access_level: str = "enabled" class _MergeRequest(BaseModel): @@ -40,6 +44,7 @@ class _MergeRequest(BaseModel): author: _User merged_at: str | None updated_at: str + created_at: datetime | None = None sha: str | None = None source_branch: str source_project_id: int | None @@ -52,7 +57,8 @@ class _MergeRequest(BaseModel): "title": self.title, "body": self.description or "", "html_url": self.web_url, - "user": {"login": self.author.username}, + "user": {"login": self.author.username, "type": "Bot" if self.author.bot else "User"}, + "created_at": self.created_at, "merged_at": self.merged_at, "updated_at": self.updated_at, "head": { @@ -94,11 +100,20 @@ class _Commit(BaseModel): message: str +class _Issue(BaseModel): + iid: int + created_at: datetime + labels: tuple[str, ...] = () + + class GitLab: def __init__(self, settings: ROISettings, transport: httpx.AsyncBaseTransport | None = None) -> None: self.settings: Final = settings token: Final = settings.gitlab_token.get_secret_value() - self.headers: Final = {"Accept": "application/json", **({"PRIVATE-TOKEN": token} if token else {})} + authorization: Final = ( + {"Authorization": f"Bearer {token}"} if settings.connection_type == "app" else {"PRIVATE-TOKEN": token} + ) + self.headers: Final = {"Accept": "application/json", **(authorization if token else {})} self.client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.ROICalculator, params={"timeout": 45, "follow_redirects": False, "transport": transport}, @@ -171,7 +186,9 @@ class GitLab: params: Final = { "simple": "true", "search": query, - **({"membership": "true"} if self.headers.get("PRIVATE-TOKEN") else {}), + **( + {"membership": "true"} if self.headers.get("PRIVATE-TOKEN") or self.headers.get("Authorization") else {} + ), } items, more = await self._page("projects", _Project, params, page) return tuple((item.path_with_namespace, item.visibility, item.archived) for item in items), more @@ -193,6 +210,8 @@ class GitLab: "state": "merged", "scope": "all", "updated_after": start.isoformat() + "T00:00:00Z", + "merged_after": start.isoformat() + "T00:00:00Z", + "merged_before": (end + timedelta(days=1)).isoformat() + "T00:00:00Z", "order_by": "updated_at", "sort": "desc", }, @@ -218,6 +237,26 @@ class GitLab: self.profiles = MappingProxyType({**self.profiles, login.casefold(): email}) return email + async def issues(self, repo: str, start: date, end: date) -> tuple[ObservedIssue, ...] | None: + project: Final = await self._project(repo) + if not project.issues_enabled or project.issues_access_level == "disabled": + return None + issues: Final = await self._all( + f"projects/{project.id}/issues", + _Issue, + { + "scope": "all", + "state": "all", + "created_after": f"{start}T00:00:00Z", + "created_before": f"{end + timedelta(days=1)}T00:00:00Z", + }, + ) + return tuple( + ObservedIssue(repo=repo, number=issue.iid, created_at=issue.created_at, labels=issue.labels) + for issue in issues + if start <= issue.created_at.date() <= end + ) + async def evidence(self, repo: str, pull: GitHubPullListItem) -> ROIPullEvidence: project: Final = await self._project(repo) path: Final = f"projects/{project.id}/merge_requests/{pull.number}" diff --git a/litellm/proxy/roi_calculator/oauth.py b/litellm/proxy/roi_calculator/oauth.py new file mode 100644 index 00000000000..32d8cd3e1d0 --- /dev/null +++ b/litellm/proxy/roi_calculator/oauth.py @@ -0,0 +1,401 @@ +import asyncio +import hashlib +import json +import os +import re +import secrets +from collections.abc import Mapping +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import Final, Literal, Protocol, TypeAlias, cast +from urllib.parse import parse_qsl, urlencode, urlsplit + +import httpx +from fastapi import HTTPException +from oauthlib.oauth2 import WebApplicationClient +from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter + +from litellm.llms.custom_httpx.http_handler import ( + get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # shared client factory has untyped params +) +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.proxy.roi_calculator.settings import ( + active_connection, + connection_id, + load_settings, + load_stored_settings, + save_settings, + select_connection, + stored_connections, +) +from litellm.proxy.roi_calculator.sync_store import SyncStore +from litellm.repositories.config_repository import ConfigRepository +from litellm.types.llms.custom_http import httpxSpecialProvider +from litellm.types.roi_calculator import ROISettings, ROISyncStatus + +Provider: TypeAlias = Literal["github", "gitlab"] +_STATE_PREFIX: Final = "roi_oauth_state_" + + +@dataclass(frozen=True, slots=True) +class OAuthConfig: + provider: Provider + api_url: str + base_url: str + client_id: str + client_secret: SecretStr + proxy_url: str + app_slug: str = "" + + @property + def cookie_path(self) -> str: + return urlsplit(self.proxy_url).path + "/roi-calculator/observed/oauth" + + @property + def installation_url(self) -> str | None: + if self.provider != "github" or not self.app_slug: + return None + path: Final = "apps" if self.api_url == "https://api.github.com" else "github-apps" + return f"{self.base_url}/{path}/{self.app_slug}/installations/new" + + @property + def redirect_uri(self) -> str: + return f"{self.proxy_url}/roi-calculator/observed/oauth/{self.provider}/callback" + + @property + def authorize_url(self) -> str: + return self.base_url + ("/login/oauth/authorize" if self.provider == "github" else "/oauth/authorize") + + @property + def token_url(self) -> str: + return self.base_url + ("/login/oauth/access_token" if self.provider == "github" else "/oauth/token") + + +def oauth_config(provider: Provider) -> OAuthConfig | None: + prefix: Final = f"LITELLM_ROI_{provider.upper()}_" + client_id: Final = os.environ.get(prefix + "CLIENT_ID", "") + client_secret: Final = os.environ.get(prefix + "CLIENT_SECRET", "") + proxy_url: Final = os.environ.get("PROXY_BASE_URL", "").rstrip("/") + if not client_id or not client_secret or not proxy_url: + return None + try: + parsed: Final = urlsplit(proxy_url) + except ValueError: + return None + if parsed.username or parsed.password or parsed.query or parsed.fragment or not parsed.hostname: + return None + if parsed.scheme != "https" and not (parsed.scheme == "http" and parsed.hostname in ("localhost", "127.0.0.1")): + return None + base: Final = os.environ.get( + prefix + "URL", "https://github.com" if provider == "github" else "https://gitlab.com" + ).rstrip("/") + api_url: Final = ( + "https://api.github.com" + if base == "https://github.com" + else base + ("/api/v3" if provider == "github" else "/api/v4") + ) + try: + validated: Final = ROISettings.model_validate( + { + "source_provider": provider, + ("github_api_url" if provider == "github" else "gitlab_api_url"): api_url, + } + ) + except ValueError: + return None + app_slug: Final = os.environ.get(prefix + "APP_SLUG", "") + if app_slug and not re.fullmatch(r"[A-Za-z0-9-]+", app_slug): + return None + return OAuthConfig( + provider, validated.source_api_url, base, client_id, SecretStr(client_secret), proxy_url, app_slug + ) + + +class OAuthState(BaseModel): + model_config = ConfigDict(frozen=True) + + provider: Provider + client_id: str + api_url: str + browser_nonce: SecretStr + verifier: SecretStr + expires_at: datetime + settings_revision: int + flow: Literal["authorize", "install"] = "authorize" + + +class _Envelope(BaseModel): + payload: str + + +class _StateRow(BaseModel): + param_value: _Envelope + + +class _Database(Protocol): + async def query_raw(self, query: str, *args: object) -> object: ... + async def execute_raw(self, query: str, *args: object) -> int: ... + + +def _state_key(state: str) -> str: + return _STATE_PREFIX + hashlib.sha256(state.encode()).hexdigest() + + +async def begin_authorization( + repository: ConfigRepository, config: OAuthConfig, *, install: bool = False +) -> tuple[str, str]: + client: Final = WebApplicationClient(config.client_id) + verifier: Final = TypeAdapter(str).validate_python(client.create_code_verifier(64)) + challenge: Final = TypeAdapter(str).validate_python(client.create_code_challenge(verifier, "S256")) + state: Final = secrets.token_urlsafe(32) + nonce: Final = secrets.token_urlsafe(32) + stored: Final = await load_stored_settings(repository) + value: Final = OAuthState( + settings_revision=stored.revision, + flow="install" if install else "authorize", + provider=config.provider, + client_id=config.client_id, + api_url=config.api_url, + browser_nonce=SecretStr(nonce), + verifier=SecretStr(verifier), + expires_at=datetime.now(timezone.utc) + timedelta(minutes=10), + ) + encoded: Final = json.dumps({**value.model_dump(mode="json"), "browser_nonce": nonce, "verifier": verifier}) + payload: Final = TypeAdapter(str).validate_python(encrypt_value_helper(encoded)) + await repository.set_param(_state_key(state), _Envelope(payload=payload).model_dump(mode="json")) + delegate: Final = repository.prisma_client.writer_db + database: Final = cast(_Database, delegate) # cast-ok: Prisma delegates database methods dynamically + await database.execute_raw( + "DELETE FROM \"LiteLLM_Config\" WHERE starts_with(param_name, $1) AND last_run_at < NOW() - INTERVAL '20 minutes'", + _STATE_PREFIX, + ) + url: Final = TypeAdapter(str).validate_python( + client.prepare_request_uri( # pyright: ignore[reportUnknownMemberType] # oauthlib leaves extension kwargs untyped + config.authorize_url, + redirect_uri=config.redirect_uri, + scope="read_api read_user" if config.provider == "gitlab" else None, + state=state, + code_challenge=challenge, + code_challenge_method="S256", + ) + ) + if install and config.installation_url: + return config.installation_url + "?" + urlencode({"state": state}), nonce + return url, nonce + + +async def consume_state( + repository: ConfigRepository, + state: str, + nonce: str, + config: OAuthConfig, + *, + flow: Literal["authorize", "install"] = "authorize", +) -> OAuthState: + if not state or not nonce or len(state) > 200: + raise HTTPException(400, "The connection expired. Start again from the ROI Calculator.") + delegate: Final = repository.prisma_client.writer_db + database: Final = cast(_Database, delegate) # cast-ok: Prisma delegates database methods dynamically + rows: Final = TypeAdapter(tuple[_StateRow, ...]).validate_python( + await database.query_raw( + 'DELETE FROM "LiteLLM_Config" WHERE param_name = $1 RETURNING param_value', + _state_key(state), + ) + ) + if not rows: + raise HTTPException(400, "This connection was already used or expired. Start again.") + plaintext: Final = decrypt_value_helper(rows[0].param_value.payload, _state_key(state)) + if plaintext is None: + raise HTTPException(400, "Could not verify the connection. Start again.") + value: Final = OAuthState.model_validate_json(plaintext) + if ( + value.flow != flow + or not secrets.compare_digest(value.browser_nonce.get_secret_value(), nonce) + or value.expires_at < datetime.now(timezone.utc) + or (value.provider, value.client_id, value.api_url) != (config.provider, config.client_id, config.api_url) + ): + raise HTTPException(400, "Could not verify the connection. Start again in the same browser.") + return value + + +class TokenGrant(BaseModel): + access_token: SecretStr + token_type: str = "bearer" + refresh_token: SecretStr = SecretStr("") + expires_in: int | None = Field(default=None, gt=0) + + +async def _token(config: OAuthConfig, body: str, transport: httpx.AsyncBaseTransport | None) -> TokenGrant: + client: Final = get_async_httpx_client( + llm_provider=httpxSpecialProvider.ROICalculator, + params={"timeout": 30, "follow_redirects": False, "transport": transport}, + ).client + try: + response: Final = await client.post( + config.token_url, data=dict(parse_qsl(body)), headers={"Accept": "application/json"} + ) + except httpx.RequestError: + raise HTTPException(502, "Could not reach the provider. Try connecting again.") from None + finally: + if transport is not None: + await client.aclose() + if response.status_code != 200: + raise HTTPException(502, "The provider rejected the connection. Try connecting again.") + try: + result: Final = TokenGrant.model_validate(response.json()) + except ValueError: + raise HTTPException(502, "The provider did not return a valid token. Try connecting again.") from None + if not result.access_token.get_secret_value() or result.token_type.casefold() != "bearer": + raise HTTPException(502, "The provider returned an unsupported token.") + return result + + +async def exchange_code( + config: OAuthConfig, state: OAuthState, code: str, transport: httpx.AsyncBaseTransport | None = None +) -> TokenGrant: + client: Final = WebApplicationClient(config.client_id) + body: Final = TypeAdapter(str).validate_python( + client.prepare_request_body( # pyright: ignore[reportUnknownMemberType] # oauthlib leaves extension kwargs untyped + code=code, + redirect_uri=config.redirect_uri, + code_verifier=state.verifier.get_secret_value(), + client_secret=config.client_secret.get_secret_value(), + ) + ) + return await _token(config, body, transport) + + +async def save_grant( + repository: ConfigRepository, + config: OAuthConfig, + grant: TokenGrant, + *, + revision: int | None = None, + previous: ROISettings | None = None, + attempt: int = 0, +) -> ROISettings: + stored: Final = await load_stored_settings(repository) + selected: Final = next( + (entry for entry in stored_connections(stored) if entry.id == connection_id(config.provider, config.api_url)), + None, + ) + scoped: Final = select_connection(stored, selected) if selected else stored + current: Final = await load_settings(repository, scoped) + if revision is not None and stored.revision != revision: + raise HTTPException(409, "The connection changed during authorization. Start again from Connections.") + if previous is not None and ( + current.source_provider, + current.source_api_url, + current.connection_type, + current.gitlab_token if current.source_provider == "gitlab" else current.github_token, + current.oauth_refresh_token, + ) != ( + previous.source_provider, + previous.source_api_url, + previous.connection_type, + previous.gitlab_token if previous.source_provider == "gitlab" else previous.github_token, + previous.oauth_refresh_token, + ): + return current + changed: Final = (current.source_provider, current.source_api_url) != (config.provider, config.api_url) + refresh_token: Final = ( + grant.refresh_token + if grant.refresh_token.get_secret_value() or previous is None + else previous.oauth_refresh_token + ) + fields: Final[Mapping[str, object]] = { + "report_mode": "observed" if previous is None else current.report_mode, + "source_provider": config.provider, + "connection_type": "app", + "repos": () if changed else current.repos, + "identity_map": {} if changed else current.identity_map, + "ignored_logins": () if changed else current.ignored_logins, + "oauth_refresh_token": refresh_token, + "oauth_expires_at": datetime.now(timezone.utc) + timedelta(seconds=grant.expires_in) + if grant.expires_in + else None, + ("github_token" if config.provider == "github" else "gitlab_token"): grant.access_token, + ("github_api_url" if config.provider == "github" else "gitlab_api_url"): config.api_url, + } + settings: Final = ROISettings.model_validate({**current.model_dump(), **fields}) + encrypted: Final = TypeAdapter(str).validate_python(encrypt_value_helper(grant.access_token.get_secret_value())) + try: + await save_settings( + repository, + settings, + encrypted if config.provider == "github" else scoped.github_token, + stored.estimator_key, + encrypted if config.provider == "gitlab" else scoped.gitlab_token, + revision=stored.revision, + ) + return settings + except HTTPException as exc: + if exc.status_code != 409 or previous is None or attempt == 2: + raise + + return await save_grant(repository, config, grant, revision=revision, previous=previous, attempt=attempt + 1) + + +def _expired(settings: ROISettings) -> bool: + return bool( + settings.connection_type == "app" + and settings.oauth_expires_at is not None + and settings.oauth_expires_at <= datetime.now(timezone.utc) + timedelta(minutes=5) + ) + + +async def connected_settings( + repository: ConfigRepository, transport: httpx.AsyncBaseTransport | None = None, selected_id: str | None = None +) -> ROISettings: + initial: Final = await load_settings(repository, selected_id=selected_id) + if not _expired(initial): + return initial + selected: Final = selected_id or active_connection(await load_stored_settings(repository)).id + store: Final = SyncStore(repository.prisma_client, "roi_oauth_refresh_" + selected) + owner: Final = secrets.token_urlsafe(24) + status: Final = ROISyncStatus( + running=True, + phase="spend", + stage="Refreshing connection", + done=0, + total=0, + estimated=0, + reused=0, + needs_attention=0, + error=None, + ) + + async def wait_for_connection(attempt: int) -> ROISettings: + if attempt >= 100: + raise HTTPException(409, "The connection is refreshing. Try again shortly.") + settings: Final = await load_settings(repository, selected_id=selected) + if not _expired(settings): + return settings + if not await store.acquire(owner, status): + await asyncio.sleep(0.1) + return await wait_for_connection(attempt + 1) + try: + current: Final = await load_settings(repository, selected_id=selected) + if not _expired(current): + return current + config: Final = oauth_config(current.source_provider) + if ( + config is None + or config.api_url != current.source_api_url + or not current.oauth_refresh_token.get_secret_value() + ): + raise HTTPException(409, "The app connection expired. Reconnect from Connections.") + client: Final = WebApplicationClient(config.client_id) + body: Final = TypeAdapter(str).validate_python( + client.prepare_refresh_body( # pyright: ignore[reportUnknownMemberType] # oauthlib leaves extension kwargs untyped + refresh_token=current.oauth_refresh_token.get_secret_value(), + client_id=config.client_id, + client_secret=config.client_secret.get_secret_value(), + ) + ) + grant: Final = await _token(config, body, transport) + return await save_grant(repository, config, grant, previous=current) + finally: + await store.finish(owner, status.model_copy(update={"running": False, "phase": "complete"})) + + return await wait_for_connection(0) diff --git a/litellm/proxy/roi_calculator/observed_analytics.py b/litellm/proxy/roi_calculator/observed_analytics.py new file mode 100644 index 00000000000..755537aadf0 --- /dev/null +++ b/litellm/proxy/roi_calculator/observed_analytics.py @@ -0,0 +1,191 @@ +import calendar +import re +from collections.abc import Mapping +from datetime import date, datetime, timedelta, timezone +from itertools import chain +from statistics import median +from types import MappingProxyType +from typing import Final + +from litellm.proxy.roi_calculator.analytics import normalize_email +from litellm.proxy.roi_calculator.branch_spend import attribute_branch_keys +from litellm.types.roi_observed import ( + ObservedAccount, + ObservedData, + ObservedHumanSummary, + ObservedPeriod, + ObservedPeriodData, + ObservedPeriods, + ObservedPerson, + ObservedPersonPeriod, + ObservedPersonPeriods, + ObservedPull, + ObservedPullPeriods, + ObservedPullResponse, + ObservedReport, + ObservedWindow, +) + + +def reporting_windows(now: datetime, days: int = 28) -> tuple[ObservedWindow, ObservedWindow, ObservedWindow]: + end: Final = now.astimezone(timezone.utc).date() - timedelta(days=1) + start: Final = end - timedelta(days=days - 1) + last_year_end: Final = date(end.year - 1, end.month, min(end.day, calendar.monthrange(end.year - 1, end.month)[1])) + return ( + ObservedWindow(start=start, end=end), + ObservedWindow(start=start - timedelta(days=days), end=start - timedelta(days=1)), + ObservedWindow(start=last_year_end - timedelta(days=days - 1), end=last_year_end), + ) + + +def declared_requester(author: str, body: str) -> str: + if author.casefold().removesuffix("[bot]") not in ("devin-ai-integration", "devin-ai"): + return "" + matches: Final = frozenset( + match.group(1) for match in re.finditer(r"^Requested by:\s*@([A-Za-z0-9_.-]+)\s*$", body, re.MULTILINE) + ) + return next(iter(matches)).casefold() if len(matches) == 1 else "" + + +def merge_hours(pull: ObservedPull) -> float | None: + if pull.created_at is None or pull.created_at.tzinfo is None or pull.merged_at.tzinfo is None: + return None + seconds: Final = (pull.merged_at - pull.created_at).total_seconds() + return seconds / 3600 if seconds >= 0 else None + + +def median_hours(pulls: tuple[ObservedPull, ...]) -> float | None: + values: Final = tuple(hours for pull in pulls if (hours := merge_hours(pull)) is not None) + return median(values) if values else None + + +def _owner_login(pull: ObservedPull) -> str: + login: Final = (pull.requester if pull.agent else pull.author).casefold() + return f"{pull.connection_id}:{login}" if pull.connection_id and login else login + + +def identity_matches(data: ObservedData, manual: Mapping[str, str], ignored: tuple[str, ...] = ()) -> Mapping[str, str]: + pulls: Final = tuple(chain(data.current.pulls, data.previous.pulls, data.last_year.pulls)) + candidates: Final = frozenset((_owner_login(pull), normalize_email(pull.profile_email)) for pull in pulls) + gateway_emails: Final = frozenset(data.gateway_emails) + profiles: Final = { + login: email + for login, email in candidates + if login not in ignored + and email in gateway_emails + and len({value for key, value in candidates if key == login and value}) == 1 + } + return MappingProxyType({**profiles, **manual}) + + +def _person_period(data: ObservedPeriodData, email: str, identities: Mapping[str, str]) -> ObservedPersonPeriod: + pulls: Final = tuple(pull for pull in data.pulls if identities.get(_owner_login(pull)) == email) + spend: Final = tuple(row for row in data.spend if row["email"] == email) + cost: Final = sum(row["spend"] for row in spend) + days: Final = (data.window.end - data.window.start).days + 1 + return ObservedPersonPeriod( + merged_prs=len(pulls), + prs_per_week=len(pulls) * 7 / days, + median_merge_hours=median_hours(pulls), + direct_authored=sum(not pull.agent for pull in pulls), + declared_agent_owned=sum(pull.agent for pull in pulls), + gateway_recorded_spend=cost, + recorded_spend_per_attributed_pr=cost / len(pulls) if spend and pulls else None, + spend_observation="records_present" if spend else "no_records", + pr_urls=tuple(pull.url for pull in pulls), + ) + + +def _person(data: ObservedData, email: str, identities: Mapping[str, str]) -> ObservedPerson: + return ObservedPerson( + name=email.split("@", 1)[0], + email=email, + logins=tuple( + sorted(frozenset(login.rsplit(":", 1)[-1] for login, address in identities.items() if address == email)) + ), + accounts=tuple( + ObservedAccount( + connection_id=login.split(":", 1)[0] if ":" in login else "", login=login.rsplit(":", 1)[-1] + ) + for login, address in identities.items() + if address == email + ), + periods=ObservedPersonPeriods( + current=_person_period(data.current, email, identities), + previous=_person_period(data.previous, email, identities), + last_year=_person_period(data.last_year, email, identities), + ), + ) + + +def _issue_count(data: ObservedPeriodData, labels: frozenset[str]) -> int | None: + if data.issues is None: + return None + return sum(bool(labels.intersection(label.casefold().strip() for label in issue.labels)) for issue in data.issues) + + +def _period(data: ObservedPeriodData, identities: Mapping[str, str]) -> ObservedPeriod: + matched: Final = tuple(pull for pull in data.pulls if _owner_login(pull) in identities) + emails: Final = frozenset(identities.values()) + spend: Final = tuple(row for row in data.spend if row["email"] in emails) + humans: Final = tuple(pull for pull in data.pulls if pull.author and not pull.agent) + return ObservedPeriod( + window=data.window, + merged_prs=len(data.pulls), + median_merge_hours=median_hours(data.pulls), + human_authored=len(humans), + agent_authored=sum(pull.agent for pull in data.pulls), + missing_author=sum(not pull.author for pull in data.pulls), + agents_without_requester=sum(pull.agent and not pull.requester for pull in data.pulls), + matched_internal_prs=len(matched), + new_bug_labeled_issues=_issue_count(data, frozenset(("bug", "kind:bug", "type::bug"))), + new_regression_labeled_issues=_issue_count( + data, frozenset(("regression", "kind:regression", "type::regression")) + ), + explicitly_titled_revert_prs=sum( + bool(re.match(r"^revert(?:\W|$)", pull.title, re.IGNORECASE)) for pull in data.pulls + ), + matched_users_recorded_spend=sum(row["spend"] for row in spend), + spend_observation="records_present" if spend else "no_records", + human_summary=ObservedHumanSummary(median_merge_hours=median_hours(humans)), + ) + + +def _pulls(data: ObservedPeriodData) -> tuple[ObservedPullResponse, ...]: + costs: Final = attribute_branch_keys( + tuple((pull.url, pull.number, pull.source_repo, pull.source_branch) for pull in data.pulls), data.branch_spend + ) + return tuple( + ObservedPullResponse.model_validate( + {**pull.model_dump(), "merge_hours": merge_hours(pull), "branch_cost": costs[(pull.url, pull.number)]} + ) + for pull in data.pulls + ) + + +def summarize_observed(data: ObservedData, manual: Mapping[str, str], ignored: tuple[str, ...] = ()) -> ObservedReport: + identities: Final = identity_matches(data, manual, ignored) + all_pulls: Final = tuple(chain(data.current.pulls, data.previous.pulls, data.last_year.pulls)) + current_pulls: Final = _pulls(data.current) + linked_branches: Final = frozenset( + (pull.source_repo, pull.source_branch) for pull in current_pulls if pull.branch_cost.status == "matched" + ) + return ObservedReport( + source_provider=data.source_provider, + connections=data.connections, + repos=data.repos, + captured_at=data.captured_at, + periods=ObservedPeriods( + current=_period(data.current, identities), + previous=_period(data.previous, identities), + last_year=_period(data.last_year, identities), + ), + people=tuple(_person(data, email, identities) for email in sorted(frozenset(identities.values()))), + pulls=ObservedPullPeriods( + current=current_pulls, previous=_pulls(data.previous), last_year=_pulls(data.last_year) + ), + unlinked_branches=tuple( + row for row in data.current.branch_spend or () if (row.repo, row.branch) not in linked_branches + ), + unmatched_logins=tuple(sorted(frozenset(_owner_login(pull) for pull in all_pulls) - identities.keys() - {""})), + ) diff --git a/litellm/proxy/roi_calculator/observed_sync.py b/litellm/proxy/roi_calculator/observed_sync.py new file mode 100644 index 00000000000..c2c5ae0e31c --- /dev/null +++ b/litellm/proxy/roi_calculator/observed_sync.py @@ -0,0 +1,265 @@ +import asyncio +from collections.abc import Awaitable, Callable, Mapping +from contextlib import suppress +from datetime import datetime, timezone +from itertools import chain +from types import MappingProxyType +from typing import Final, TypeAlias +from uuid import uuid4 + +import httpx + +from litellm._logging import verbose_proxy_logger +from litellm.proxy.roi_calculator.analytics import normalize_email +from litellm.proxy.roi_calculator.github import GitHub, GitHubPullListItem, SourceError +from litellm.proxy.roi_calculator.github_observed import GitHubObserved +from litellm.proxy.roi_calculator.gitlab import GitLab +from litellm.proxy.roi_calculator.observed_analytics import declared_requester, reporting_windows +from litellm.proxy.roi_calculator.source import repository_tag +from litellm.proxy.roi_calculator.sync import BranchSpendReader, GatewayUserReader, SpendReader +from litellm.proxy.roi_calculator.sync_store import SyncStore +from litellm.types.roi_calculator import ROISettings, ROISyncStatus +from litellm.types.roi_observed import ObservedData, ObservedIssue, ObservedPeriodData, ObservedPull, ObservedWindow + + +def _author(pull: GitHubPullListItem) -> str: + return pull.user.login or "" if pull.user else "" + + +def _agent(pull: GitHubPullListItem) -> bool: + login: Final = _author(pull) + return ( + bool(pull.user and pull.user.type == "Bot") + or login.endswith("[bot]") + or bool(login.startswith(("project_", "group_")) and "_bot_" in login) + ) + + +def _owner(pull: GitHubPullListItem) -> str: + login: Final = _author(pull) + return declared_requester(login, pull.body or "") if _agent(pull) else login + + +def _public_email(pull: GitHubPullListItem) -> str: + return normalize_email(pull.user.email) if pull.user and not _agent(pull) else "" + + +def _pull(settings: ROISettings, repo: str, pull: GitHubPullListItem, profiles: Mapping[str, str]) -> ObservedPull: + if not pull.merged_at: + raise SourceError("The source returned an unmerged change. No partial report was saved.") + return ObservedPull( + repo=repo, + number=pull.number, + title=pull.title, + url=pull.html_url, + author=_author(pull), + agent=_agent(pull), + requester=declared_requester(_author(pull), pull.body or "") if _agent(pull) else "", + profile_email=profiles.get(_owner(pull).casefold(), "") or _public_email(pull), + created_at=pull.created_at, + merged_at=datetime.fromisoformat(pull.merged_at.replace("Z", "+00:00")), + source_repo=repository_tag(settings, pull.head.repo.full_name) if pull.head and pull.head.repo else "", + source_branch=pull.head.ref if pull.head else "", + ) + + +async def collect_observed( + settings: ROISettings, + spend_reader: SpendReader, + gateway_user_reader: GatewayUserReader, + branch_spend_reader: BranchSpendReader, + now: datetime, + progress: Callable[[str, int, int], None], + transport: httpx.AsyncBaseTransport | None = None, + days: int = 28, +) -> ObservedData: + source: Final = GitLab(settings, transport) if settings.source_provider == "gitlab" else GitHub(settings, transport) + activity: Final = ( + GitHubObserved(settings, source.client) + if settings.source_provider == "github" and settings.github_token.get_secret_value() + else source + ) + windows: Final = reporting_windows(now.astimezone(timezone.utc), days) + slots: Final = asyncio.Semaphore(4) + total: Final = len(settings.repos) * 3 + + async def profile(login: str) -> tuple[str, str]: + async with slots: + return login.casefold(), await source.profile_email(login) + + async def repository( + repo: str, window: ObservedWindow + ) -> tuple[tuple[ObservedPull, ...], tuple[ObservedIssue, ...] | None]: + raw: Final = await activity.pulls(repo, window.start, window.end) + unique: Final = {(repo, item.number): item for item in raw} + if len(unique) != len(raw): + raise SourceError("The source returned duplicate changes. Retry to get a complete report.") + owners: Final = ( + frozenset(_owner(pull) for pull in raw if not _public_email(pull)) + - {""} + - settings.identity_map.keys() + - frozenset(settings.ignored_logins) + ) + profiles: Final = MappingProxyType(dict(await asyncio.gather(*(profile(login) for login in owners)))) + pulls: Final = tuple(_pull(settings, repo, item, profiles) for item in raw) + issues: Final = await activity.issues(repo, window.start, window.end) + return pulls, issues + + async def period(window: ObservedWindow, offset: int) -> ObservedPeriodData: + async def read(index: int, repo: str) -> tuple[tuple[ObservedPull, ...], tuple[ObservedIssue, ...] | None]: + progress(f"Reading {repo} ({window.start} to {window.end})", offset + index, total) + return await repository(repo, window) + + results: Final = tuple([await read(index, repo) for index, repo in enumerate(settings.repos)]) + pulls: Final = tuple(chain.from_iterable(result[0] for result in results)) + issues: Final = ( + None + if results and all(result[1] is None for result in results) + else tuple(chain.from_iterable(result[1] or () for result in results)) + ) + branches: Final = tuple( + sorted( + frozenset( + ( + *(repository_tag(settings, repo) for repo in settings.repos), + *(pull.source_repo for pull in pulls), + ) + ) + - {""} + ) + ) + return ObservedPeriodData( + window=window, + pulls=tuple(sorted(pulls, key=lambda pull: (pull.merged_at, pull.repo, pull.number), reverse=True)), + issues=issues, + spend=await spend_reader(window.start, window.end), + branch_spend=await branch_spend_reader(window.start, window.end, branches), + ) + + try: + gateway_emails: Final = await gateway_user_reader() + current: Final = await period(windows[0], 0) + previous: Final = await period(windows[1], len(settings.repos)) + last_year: Final = await period(windows[2], len(settings.repos) * 2) + progress("Saving report", total, total) + return ObservedData( + source_provider=settings.source_provider, + source_api_url=settings.source_api_url, + repos=settings.repos, + captured_at=now, + gateway_emails=tuple(sorted(gateway_emails)), + current=current, + previous=previous, + last_year=last_year, + ) + finally: + await source.close() + + +Progress: TypeAlias = Callable[[str, int, int], None] +BuildReport: TypeAlias = Callable[[Progress], Awaitable[ObservedData]] + + +class ObservedSyncManager: + def __init__(self) -> None: + self.status: ROISyncStatus = ROISyncStatus( + running=False, + phase="idle", + stage="Not synced", + done=0, + total=0, + estimated=0, + reused=0, + needs_attention=0, + error=None, + ) + self._task: asyncio.Task[None] | None = None + self._lock: Final = asyncio.Lock() + + def _progress(self, stage: str, done: int, total: int) -> None: + self.status = self.status.model_copy( + update={"phase": "repositories", "stage": stage, "done": done, "total": total} + ) + + async def start(self, build: BuildReport, store: SyncStore, scheduled_interval: float = 0) -> bool: + async with self._lock: + if self._task is not None and not self._task.done(): + return False + status: Final = ROISyncStatus( + running=True, + phase="repositories", + stage="Reading repository activity", + done=0, + total=0, + estimated=0, + reused=0, + needs_attention=0, + error=None, + started_at=datetime.now(timezone.utc).isoformat(), + ) + owner: Final = str(uuid4()) + if not await store.acquire(owner, status, scheduled_interval): + return False + self.status = status + self._task = asyncio.create_task(self._run(build, store, owner)) + return True + + async def cancel(self) -> None: + if self._task is not None and not self._task.done(): + self._task.cancel() + with suppress(asyncio.CancelledError): + await self._task + + async def _heartbeat(self, store: SyncStore, owner: str, task: asyncio.Task[object] | None) -> None: + if task is None: + return + try: + while True: + await asyncio.sleep(5) + if not await store.heartbeat(owner, self.status): + task.cancel() + return + except Exception: # noqa: BLE001 # loss of the database lease must stop publication + task.cancel() + + async def _run(self, build: BuildReport, store: SyncStore, owner: str) -> None: + monitor: Final = asyncio.create_task(self._heartbeat(store, owner, asyncio.current_task())) + try: + report: Final = await build(self._progress) + monitor.cancel() + with suppress(asyncio.CancelledError): + await monitor + complete: Final = self.status.model_copy( + update={ + "running": False, + "phase": "complete", + "stage": "Up to date", + "finished_at": datetime.now(timezone.utc).isoformat(), + } + ) + if not await store.finish(owner, complete, report): + raise SourceError("The sync was cancelled or replaced. The previous report was kept.") + self.status = complete + except asyncio.CancelledError: + self.status = self.status.model_copy(update={"phase": "cancelled", "stage": "Sync cancelled"}) + raise + except SourceError as exc: + self.status = self.status.model_copy(update={"phase": "error", "stage": "Sync failed", "error": str(exc)}) + except Exception: # noqa: BLE001 # background tasks must persist a safe error without exposing credentials + verbose_proxy_logger.exception("Observed ROI sync failed") + self.status = self.status.model_copy( + update={ + "phase": "error", + "stage": "Sync failed", + "error": "Could not finish syncing. The previous report was kept. Retry after checking the connection.", + } + ) + finally: + monitor.cancel() + with suppress(asyncio.CancelledError): + await monitor + self.status = self.status.model_copy( + update={"running": False, "finished_at": datetime.now(timezone.utc).isoformat()} + ) + if self.status.phase != "complete": + await store.finish(owner, self.status) diff --git a/litellm/proxy/roi_calculator/observed_workspace.py b/litellm/proxy/roi_calculator/observed_workspace.py new file mode 100644 index 00000000000..f08e0f5d681 --- /dev/null +++ b/litellm/proxy/roi_calculator/observed_workspace.py @@ -0,0 +1,135 @@ +from collections.abc import Mapping +from datetime import date, datetime +from itertools import chain +from types import MappingProxyType +from typing import Final + +import httpx + +from litellm.proxy.roi_calculator.github import SourceError +from litellm.proxy.roi_calculator.observed_analytics import reporting_windows, summarize_observed +from litellm.proxy.roi_calculator.observed_sync import Progress, collect_observed +from litellm.proxy.roi_calculator.settings import StoredConnection, connection_id +from litellm.proxy.roi_calculator.source import repository_tag +from litellm.proxy.roi_calculator.sync import BranchSpendReader, GatewayUserReader, SpendReader +from litellm.types.roi_calculator import ROISettings, ROISpendRecord +from litellm.types.roi_observed import ObservedData, ObservedPeriodData, ObservedReport, ObservedSource + + +def source_details(settings: ROISettings) -> ObservedSource: + return ObservedSource( + id=connection_id(settings.source_provider, settings.source_api_url), + source_provider=settings.source_provider, + api_url=settings.source_api_url, + repos=settings.repos, + ) + + +def scoped_data(data: ObservedData, source: ObservedSource) -> ObservedData: + def period(value: ObservedPeriodData) -> ObservedPeriodData: + return value.model_copy( + update={"pulls": tuple(pull.model_copy(update={"connection_id": source.id}) for pull in value.pulls)} + ) + + return data.model_copy( + update={ + "connections": (source,), + "current": period(data.current), + "previous": period(data.previous), + "last_year": period(data.last_year), + } + ) + + +def summarize_workspace(data: ObservedData, connections: tuple[StoredConnection, ...]) -> ObservedReport: + included: Final = frozenset(source.id for source in data.connections) + active: Final = tuple(entry for entry in connections if entry.id in included) + maps: Final = ({f"{entry.id}:{login}": email for login, email in entry.identity_map.items()} for entry in active) + identities: Final = MappingProxyType(dict(chain.from_iterable(mapping.items() for mapping in maps))) + ignored: Final = tuple( + chain.from_iterable(tuple(f"{entry.id}:{login}" for login in entry.ignored_logins) for entry in connections) + ) + return summarize_observed(data, identities, ignored) + + +def combine_observed(sources: tuple[ObservedData, ...], repos: tuple[str, ...]) -> ObservedData: + first: Final = sources[0] + + def period(values: tuple[ObservedPeriodData, ...]) -> ObservedPeriodData: + if any(value.window != values[0].window or value.spend != values[0].spend for value in values): + raise SourceError("The reporting windows changed during sync. Retry to get a complete report.") + pulls: Final = tuple(chain.from_iterable(value.pulls for value in values)) + if len({pull.url for pull in pulls}) != len(pulls): + raise SourceError("A repository is selected through more than one connection. Select it once.") + return ObservedPeriodData( + window=values[0].window, + pulls=tuple(sorted(pulls, key=lambda pull: (pull.merged_at, pull.url), reverse=True)), + issues=None + if all(value.issues is None for value in values) + else tuple(chain.from_iterable(value.issues or () for value in values)), + spend=values[0].spend, + branch_spend=None + if any(value.branch_spend is None for value in values) + else tuple( + { + (row.repo, row.branch): row + for row in chain.from_iterable(value.branch_spend or () for value in values) + }.values() + ), + ) + + providers: Final = frozenset(source.source_provider for source in sources) + return ObservedData( + source_provider=first.source_provider if len(providers) == 1 else "mixed", + source_api_url=first.source_api_url if len(sources) == 1 else "", + connections=tuple(chain.from_iterable(source.connections for source in sources)), + repos=repos, + captured_at=first.captured_at, + gateway_emails=first.gateway_emails, + current=period(tuple(source.current for source in sources)), + previous=period(tuple(source.previous for source in sources)), + last_year=period(tuple(source.last_year for source in sources)), + ) + + +async def collect_workspace( + connections: tuple[tuple[ROISettings, BranchSpendReader], ...], + spend_reader: SpendReader, + gateway_user_reader: GatewayUserReader, + now: datetime, + progress: Progress, + transport: httpx.AsyncBaseTransport | None = None, + days: int = 28, +) -> ObservedData: + windows: Final = reporting_windows(now, days) + spending: Final[Mapping[tuple[date, date], tuple[ROISpendRecord, ...]]] = { + (window.start, window.end): await spend_reader(window.start, window.end) for window in windows + } + emails: Final = await gateway_user_reader() + total: Final = sum(len(settings.repos) * 3 for settings, _reader in connections) + + async def spend(start: date, end: date) -> tuple[ROISpendRecord, ...]: + return spending[(start, end)] + + async def users() -> frozenset[str]: + return emails + + async def collect(index: int, settings: ROISettings, branch_reader: BranchSpendReader) -> ObservedData: + offset: Final = sum(len(prior.repos) * 3 for prior, _reader in connections[:index]) + + def update(stage: str, done: int, _total: int) -> None: + progress(stage, offset + done, total) + + data: Final = await collect_observed(settings, spend, users, branch_reader, now, update, transport, days=days) + return scoped_data(data, source_details(settings)) + + data: Final = tuple( + [await collect(index, settings, reader) for index, (settings, reader) in enumerate(connections)] + ) + repos: Final = tuple( + chain.from_iterable( + tuple(repository_tag(settings, repo) if len(connections) > 1 else repo for repo in settings.repos) + for settings, _reader in connections + ) + ) + return combine_observed(data, repos) diff --git a/litellm/proxy/roi_calculator/settings.py b/litellm/proxy/roi_calculator/settings.py new file mode 100644 index 00000000000..89a691843f4 --- /dev/null +++ b/litellm/proxy/roi_calculator/settings.py @@ -0,0 +1,253 @@ +from collections.abc import Mapping +from datetime import datetime +from hashlib import sha256 +from types import MappingProxyType +from typing import Annotated, Final, Literal + +from fastapi import Depends, HTTPException +from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError + +from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper +from litellm.repositories.config_repository import ConfigRepository +from litellm.types.roi_calculator import DEFAULT_PROMPT, ROISettings + +_SETTINGS_KEY: Final = "roi_calculator_settings" + + +def connection_id(provider: str, api_url: str) -> str: + return provider + "_" + sha256(api_url.strip().rstrip("/").encode()).hexdigest()[:16] + + +class StoredConnection(BaseModel): + model_config = ConfigDict(frozen=True) + + source_provider: Literal["github", "gitlab"] + api_url: str + token: str = "" + connection_type: Literal["token", "app"] = "token" + oauth_refresh_token: str = "" + oauth_expires_at: datetime | None = None + repos: tuple[str, ...] = () + identity_map: Mapping[str, str] = Field(default_factory=lambda: MappingProxyType({})) + ignored_logins: tuple[str, ...] = () + + @property + def id(self) -> str: + return connection_id(self.source_provider, self.api_url) + + +class StoredROISettings(BaseModel): + model_config = ConfigDict(extra="ignore") + + revision: int = 0 + report_mode: Literal["legacy", "observed"] = "legacy" + source_provider: Literal["github", "gitlab"] = "github" + connection_type: Literal["token", "app"] = "token" + oauth_refresh_token: str = "" + oauth_expires_at: datetime | None = None + ignored_logins: tuple[str, ...] = () + gitlab_api_url: str = "https://gitlab.com/api/v4" + gitlab_token: str = "" + github_api_url: str = "https://api.github.com" + github_token: str = "" + estimator_key: str = "" + repos: tuple[str, ...] = () + estimator_model: str = "" + estimator_prompt: str = DEFAULT_PROMPT + backfill_days: int = Field(default=7, ge=1, le=3650) + update_interval_minutes: float = Field(default=1440, ge=0, le=43200) + identity_map: Mapping[str, str] = Field(default_factory=lambda: MappingProxyType({})) + connections: tuple[StoredConnection, ...] = () + + +def active_connection(stored: StoredROISettings) -> StoredConnection: + return StoredConnection( + source_provider=stored.source_provider, + api_url=stored.gitlab_api_url if stored.source_provider == "gitlab" else stored.github_api_url, + token=stored.gitlab_token if stored.source_provider == "gitlab" else stored.github_token, + connection_type=stored.connection_type, + oauth_refresh_token=stored.oauth_refresh_token, + oauth_expires_at=stored.oauth_expires_at, + repos=stored.repos, + identity_map=stored.identity_map, + ignored_logins=stored.ignored_logins, + ) + + +def stored_connections(stored: StoredROISettings) -> tuple[StoredConnection, ...]: + active: Final = active_connection(stored) + if not stored.connections and not active.repos and not active.token: + return () + return tuple({**{entry.id: entry for entry in stored.connections}, active.id: active}.values()) + + +def select_connection(stored: StoredROISettings, selected: StoredConnection) -> StoredROISettings: + return stored.model_copy( + update={ + "source_provider": selected.source_provider, + "connection_type": selected.connection_type, + "oauth_refresh_token": selected.oauth_refresh_token, + "oauth_expires_at": selected.oauth_expires_at, + "repos": selected.repos, + "identity_map": selected.identity_map, + "ignored_logins": selected.ignored_logins, + ("gitlab_api_url" if selected.source_provider == "gitlab" else "github_api_url"): selected.api_url, + ("gitlab_token" if selected.source_provider == "gitlab" else "github_token"): selected.token, + } + ) + + +async def read_admin( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> UserAPIKeyAuth: + if user_api_key_dict.user_role not in ( + LitellmUserRoles.PROXY_ADMIN, + LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, + ): + raise HTTPException(status_code=403, detail="Only proxy admins can access the ROI Calculator.") + return user_api_key_dict + + +async def write_admin( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], +) -> UserAPIKeyAuth: + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException(status_code=403, detail="Only proxy admins can change ROI Calculator settings.") + return user_api_key_dict + + +async def get_roi_config_repository( + _user: Annotated[UserAPIKeyAuth, Depends(read_admin)], +) -> ConfigRepository: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException( + status_code=500, + detail=CommonProxyErrors.db_not_connected_error.value, + ) + return ConfigRepository(prisma_client, use_writer=True) + + +async def load_stored_settings(repository: ConfigRepository, selected_id: str | None = None) -> StoredROISettings: + parameter: Final = await repository.get_param(_SETTINGS_KEY) + if parameter is None: + if selected_id is not None: + raise HTTPException(404, "This connection no longer exists. Reload Connections.") + return StoredROISettings() + try: + stored: Final = StoredROISettings.model_validate(parameter.param_value) + except ValidationError: + raise HTTPException(status_code=500, detail="Stored ROI Calculator settings are invalid.") from None + if selected_id is None: + return stored + selected: Final = next((entry for entry in stored_connections(stored) if entry.id == selected_id), None) + if selected is None: + raise HTTPException(404, "This connection no longer exists. Reload Connections.") + return select_connection(stored, selected) + + +async def load_settings( + repository: ConfigRepository, value: StoredROISettings | None = None, selected_id: str | None = None +) -> ROISettings: + stored: Final = value if value is not None else await load_stored_settings(repository, selected_id) + token: Final = decrypt_value_helper(stored.github_token, _SETTINGS_KEY) if stored.github_token else "" + try: + return ROISettings( + report_mode=stored.report_mode, + source_provider=stored.source_provider, + connection_type=stored.connection_type, + oauth_refresh_token=SecretStr(decrypt_value_helper(stored.oauth_refresh_token, _SETTINGS_KEY) or "") + if stored.oauth_refresh_token + else SecretStr(""), + oauth_expires_at=stored.oauth_expires_at, + ignored_logins=stored.ignored_logins, + gitlab_api_url=stored.gitlab_api_url, + gitlab_token=SecretStr(decrypt_value_helper(stored.gitlab_token, _SETTINGS_KEY) or "") + if stored.gitlab_token + else SecretStr(""), + github_api_url=stored.github_api_url, + github_token=SecretStr(token or ""), + estimator_key=SecretStr(decrypt_value_helper(stored.estimator_key, _SETTINGS_KEY) or "") + if stored.estimator_key + else SecretStr(""), + update_interval_minutes=stored.update_interval_minutes, + repos=stored.repos, + estimator_model=stored.estimator_model, + estimator_prompt=stored.estimator_prompt, + backfill_days=stored.backfill_days, + identity_map=stored.identity_map, + ) + except ValidationError: + raise HTTPException(status_code=500, detail="Stored ROI Calculator settings are invalid.") from None + + +async def save_settings( + repository: ConfigRepository, + settings: ROISettings, + encrypted_token: str, + encrypted_estimator_key: str, + encrypted_gitlab_token: str = "", + revision: int = 0, + replace_connection_id: str | None = None, +) -> None: + previous: Final = await load_stored_settings(repository) + stored: Final = StoredROISettings( + revision=revision + 1, + report_mode=settings.report_mode, + source_provider=settings.source_provider, + connection_type=settings.connection_type, + oauth_refresh_token=TypeAdapter(str).validate_python( + encrypt_value_helper(settings.oauth_refresh_token.get_secret_value()) + ) + if settings.oauth_refresh_token.get_secret_value() + else "", + oauth_expires_at=settings.oauth_expires_at, + ignored_logins=settings.ignored_logins, + gitlab_api_url=settings.gitlab_api_url, + gitlab_token=encrypted_gitlab_token, + github_api_url=settings.github_api_url, + github_token=encrypted_token, + estimator_key=encrypted_estimator_key, + update_interval_minutes=settings.update_interval_minutes, + repos=settings.repos, + estimator_model=settings.estimator_model, + estimator_prompt=settings.estimator_prompt, + backfill_days=settings.backfill_days, + identity_map=settings.identity_map, + ) + active: Final = active_connection(stored) + combined: Final = stored.model_copy( + update={ + "connections": tuple( + { + **{entry.id: entry for entry in stored_connections(previous) if entry.id != replace_connection_id}, + active.id: active, + }.values() + ) + } + ) + if not await repository.set_param_if_revision(_SETTINGS_KEY, combined.model_dump(mode="json"), revision): + raise HTTPException(409, "Settings changed while you were editing. Reload and try again.") + + +async def enable_observed_reporting(repository: ConfigRepository) -> None: + stored: Final = await load_stored_settings(repository) + if stored.report_mode == "observed": + return + updated: Final = stored.model_copy(update={"report_mode": "observed", "revision": stored.revision + 1}) + if not await repository.set_param_if_revision(_SETTINGS_KEY, updated.model_dump(mode="json"), stored.revision): + raise HTTPException(409, "Settings changed while starting the report. Reload and try again.") + + +async def save_connection_identities( + repository: ConfigRepository, stored: StoredROISettings, connections: tuple[StoredConnection, ...] +) -> None: + active: Final = next(entry for entry in connections if entry.id == active_connection(stored).id) + updated: Final = select_connection(stored, active).model_copy( + update={"connections": connections, "revision": stored.revision + 1} + ) + if not await repository.set_param_if_revision(_SETTINGS_KEY, updated.model_dump(mode="json"), stored.revision): + raise HTTPException(409, "Settings changed while you were editing. Reload and try again.") diff --git a/litellm/proxy/roi_calculator/sync_store.py b/litellm/proxy/roi_calculator/sync_store.py index 43a2533eb59..df126da6d70 100644 --- a/litellm/proxy/roi_calculator/sync_store.py +++ b/litellm/proxy/roi_calculator/sync_store.py @@ -1,11 +1,12 @@ from datetime import datetime, timezone from types import MappingProxyType -from typing import Final, Protocol, cast # noqa: TID251 - PrismaWrapper dynamically delegates database methods +from typing import Final, Literal, Protocol, cast # noqa: TID251 - PrismaWrapper dynamically delegates database methods from pydantic import BaseModel, ConfigDict, TypeAdapter from litellm.proxy.utils import PrismaClient from litellm.types.roi_calculator import ROIReport, ROISyncStatus +from litellm.types.roi_observed import ObservedData _SYNC_KEY: Final = "roi_calculator_sync" _REPORT_KEY: Final = "roi_calculator_report" @@ -30,8 +31,14 @@ class _SyncDatabase(Protocol): class SyncStore: - def __init__(self, prisma: PrismaClient) -> None: + def __init__( + self, + prisma: PrismaClient, + namespace: Literal["roi_calculator", "roi_observed", "roi_oauth_refresh"] = "roi_calculator", + ) -> None: self._db: Final = cast(_SyncDatabase, prisma.writer_db) # cast-ok: PrismaWrapper delegates methods dynamically + self._sync_key: Final = namespace + "_sync" + self._report_key: Final = namespace + "_report" async def acquire(self, owner: str, status: ROISyncStatus, scheduled_interval: float = 0) -> bool: rows: Final = await self._db.query_raw( @@ -43,7 +50,7 @@ class SyncStore: OR "LiteLLM_Config".param_value->'status'->>'running' = 'false') AND ($3::text::double precision = 0 OR "LiteLLM_Config".last_run_at <= NOW() - $3::text::double precision * INTERVAL '1 minute') RETURNING param_name""", - _SYNC_KEY, + self._sync_key, _SyncState(owner=owner, status=status).model_dump_json(), str(scheduled_interval), ) @@ -58,14 +65,20 @@ class SyncStore: AND param_value->'status'->>'running' = 'true' AND last_run_at >= NOW() - INTERVAL '60 seconds' RETURNING param_name""", - _SYNC_KEY, + self._sync_key, owner, status.model_dump_json(), ) return bool(rows) - async def finish(self, owner: str, status: ROISyncStatus, report: ROIReport | None = None) -> bool: - report_json: Final = TypeAdapter(ROIReport).dump_json(report).decode() if report is not None else None + async def finish(self, owner: str, status: ROISyncStatus, report: ROIReport | ObservedData | None = None) -> bool: + report_json: Final = ( + report.model_dump_json() + if isinstance(report, ObservedData) + else TypeAdapter(ROIReport).dump_json(report).decode() + if report is not None + else None + ) rows: Final = await self._db.query_raw( """WITH owned AS ( SELECT param_name FROM "LiteLLM_Config" @@ -92,11 +105,11 @@ class SyncStore: UPDATE "LiteLLM_Config" SET param_value = jsonb_set(param_value, '{status}', $3::jsonb), last_run_at = NOW() WHERE param_name IN (SELECT param_name FROM owned) RETURNING param_name""", - _SYNC_KEY, + self._sync_key, owner, status.model_dump_json(), report_json, - _REPORT_KEY, + self._report_key, ) return bool(rows) @@ -105,7 +118,7 @@ class SyncStore: await self._db.query_raw( """SELECT param_value, last_run_at, last_run_at < NOW() - INTERVAL '60 seconds' AS expired FROM "LiteLLM_Config" WHERE param_name = $1""", - _SYNC_KEY, + self._sync_key, ) ) if not rows: @@ -119,7 +132,7 @@ class SyncStore: "phase": "error", "finished_at": rows[0].last_run_at.replace(tzinfo=timezone.utc).isoformat(), "stage": "Sync interrupted", - "error": "The worker stopped responding. Run analysis again to resume saved estimates.", + "error": "The worker stopped responding. Sync again to refresh the report.", } ) ) @@ -136,8 +149,8 @@ class SyncStore: ) ), last_run_at = NOW() WHERE param_name = $1 AND param_value->'status'->>'running' = 'true' """, - _SYNC_KEY, + self._sync_key, ) async def clear_report(self) -> None: - await self._db.execute_raw('DELETE FROM "LiteLLM_Config" WHERE param_name = $1', _REPORT_KEY) + await self._db.execute_raw('DELETE FROM "LiteLLM_Config" WHERE param_name = $1', self._report_key) diff --git a/litellm/repositories/config_repository.py b/litellm/repositories/config_repository.py index c5674a4b398..c0be8ac55a9 100644 --- a/litellm/repositories/config_repository.py +++ b/litellm/repositories/config_repository.py @@ -33,6 +33,10 @@ class _ConfigTable(Protocol): async def delete(self, *, where: Mapping[str, str]) -> _ConfigRow | None: ... +class _ConfigDatabase(Protocol): + async def query_raw(self, query: str, *args: object) -> object: ... + + class ConfigParam: """Simple wrapper for config parameter from DB.""" @@ -85,6 +89,26 @@ class ConfigRepository: ) return ConfigParam(param_name=param_name, param_value=param_value) + async def set_param_if_revision(self, param_name: str, param_value: object, revision: int) -> bool: + delegate: Final = self.prisma_client.writer_db + database: Final = cast(_ConfigDatabase, delegate) # cast-ok: Prisma delegates database methods dynamically + rows: Final = await database.query_raw( + """INSERT INTO "LiteLLM_Config" (param_name, param_value, last_run_at) + SELECT $1, $2::jsonb, NOW() WHERE $3::int = 0 + ON CONFLICT (param_name) DO UPDATE + SET param_value = EXCLUDED.param_value, last_run_at = NOW() + WHERE COALESCE(("LiteLLM_Config".param_value->>'revision')::int, 0) = $3::int + RETURNING param_name""" + if revision == 0 + else """UPDATE "LiteLLM_Config" SET param_value = $2::jsonb, last_run_at = NOW() + WHERE param_name = $1 AND COALESCE((param_value->>'revision')::int, 0) = $3::int + RETURNING param_name""", + param_name, + json.dumps(param_value), + revision, + ) + return bool(rows) + async def delete_param(self, param_name: str) -> bool: """Delete a config parameter from the database.""" try: diff --git a/litellm/types/roi_calculator.py b/litellm/types/roi_calculator.py index 63a28ec71ca..7a6cd8ab2e1 100644 --- a/litellm/types/roi_calculator.py +++ b/litellm/types/roi_calculator.py @@ -1,4 +1,5 @@ from collections.abc import Mapping +from datetime import datetime from types import MappingProxyType from typing import Final, Literal @@ -24,7 +25,12 @@ def normalize_source_login(value: str, provider: str = "github") -> str: class ROISettings(BaseModel): model_config = ConfigDict(frozen=True) + report_mode: Literal["legacy", "observed"] = "legacy" source_provider: Literal["github", "gitlab"] = "github" + connection_type: Literal["token", "app"] = "token" + oauth_refresh_token: SecretStr = SecretStr("") + oauth_expires_at: datetime | None = None + ignored_logins: tuple[str, ...] = () gitlab_api_url: str = "https://gitlab.com/api/v4" gitlab_token: SecretStr = SecretStr("") github_api_url: str = "https://api.github.com" @@ -74,8 +80,13 @@ class ROISettings(BaseModel): import re normalized_values: Final = tuple(repo.strip().rstrip("/").removesuffix(".git") for repo in values) + repository_keys: Final = tuple( + repo.casefold() if info.data.get("source_provider") != "gitlab" else repo for repo in normalized_values + ) normalized: Final = tuple( - repo for index, repo in enumerate(normalized_values) if repo not in normalized_values[:index] + repo + for index, repo in enumerate(normalized_values) + if repository_keys[index] not in repository_keys[:index] ) pattern: Final = ( r"[A-Za-z0-9_.-]+(?:/[A-Za-z0-9_.-]+)+" @@ -121,6 +132,7 @@ class ROISettings(BaseModel): class ROISettingsUpdate(BaseModel): model_config = ConfigDict(extra="forbid") + report_mode: Literal["legacy", "observed"] | None = None source_provider: Literal["github", "gitlab"] | None = None gitlab_api_url: str | None = None gitlab_token: str | None = None @@ -140,6 +152,7 @@ class ROIEstimatorModel(BaseModel): class ROISettingsResponse(BaseModel): + report_mode: Literal["legacy", "observed"] = "legacy" source_provider: Literal["github", "gitlab"] = "github" gitlab_api_url: str = "https://gitlab.com/api/v4" has_gitlab_token: bool = False diff --git a/litellm/types/roi_observed.py b/litellm/types/roi_observed.py new file mode 100644 index 00000000000..b617175788b --- /dev/null +++ b/litellm/types/roi_observed.py @@ -0,0 +1,213 @@ +from collections.abc import Mapping +from datetime import date, datetime +from typing import Literal + +from pydantic import BaseModel, ConfigDict, Field + +from litellm.types.roi_calculator import ROIBranchAttribution, ROIBranchSpend, ROISpendRecord + + +class ObservedModel(BaseModel): + model_config = ConfigDict(frozen=True) + + +class ObservedWindow(ObservedModel): + start: date + end: date + + +class ObservedSource(ObservedModel): + id: str + source_provider: Literal["github", "gitlab"] + api_url: str + repos: tuple[str, ...] + + +class ObservedAccount(ObservedModel): + connection_id: str + login: str + + +class ObservedPull(ObservedModel): + connection_id: str = "" + repo: str + number: int + title: str + url: str + author: str + agent: bool = False + requester: str = "" + profile_email: str = "" + created_at: datetime | None = None + merged_at: datetime + source_repo: str = "" + source_branch: str = "" + + +class ObservedIssue(ObservedModel): + repo: str + number: int + created_at: datetime + labels: tuple[str, ...] = () + + +class ObservedPeriodData(ObservedModel): + window: ObservedWindow + pulls: tuple[ObservedPull, ...] + issues: tuple[ObservedIssue, ...] | None + spend: tuple[ROISpendRecord, ...] + branch_spend: tuple[ROIBranchSpend, ...] | None = None + + +class ObservedData(ObservedModel): + source_provider: Literal["github", "gitlab", "mixed"] + connections: tuple[ObservedSource, ...] = () + source_api_url: str + repos: tuple[str, ...] + captured_at: datetime + gateway_emails: tuple[str, ...] + current: ObservedPeriodData + previous: ObservedPeriodData + last_year: ObservedPeriodData + + +class ObservedPersonPeriod(ObservedModel): + merged_prs: int + prs_per_week: float + median_merge_hours: float | None + direct_authored: int + declared_agent_owned: int + gateway_recorded_spend: float + recorded_spend_per_attributed_pr: float | None + spend_observation: Literal["records_present", "no_records"] + pr_urls: tuple[str, ...] + + +class ObservedPersonPeriods(ObservedModel): + current: ObservedPersonPeriod + previous: ObservedPersonPeriod + last_year: ObservedPersonPeriod + + +class ObservedPerson(ObservedModel): + name: str + email: str + logins: tuple[str, ...] + accounts: tuple[ObservedAccount, ...] = () + periods: ObservedPersonPeriods + + +class ObservedHumanSummary(ObservedModel): + median_merge_hours: float | None + + +class ObservedPeriod(ObservedModel): + window: ObservedWindow + merged_prs: int + median_merge_hours: float | None + human_authored: int + agent_authored: int + missing_author: int + agents_without_requester: int + matched_internal_prs: int + new_bug_labeled_issues: int | None + new_regression_labeled_issues: int | None + explicitly_titled_revert_prs: int + matched_users_recorded_spend: float + spend_observation: Literal["records_present", "no_records"] + human_summary: ObservedHumanSummary + + +class ObservedPeriods(ObservedModel): + current: ObservedPeriod + previous: ObservedPeriod + last_year: ObservedPeriod + + +class ObservedPullResponse(ObservedPull): + merge_hours: float | None + branch_cost: ROIBranchAttribution + + +class ObservedPullPeriods(ObservedModel): + current: tuple[ObservedPullResponse, ...] + previous: tuple[ObservedPullResponse, ...] + last_year: tuple[ObservedPullResponse, ...] + + +class ObservedReport(ObservedModel): + source_provider: Literal["github", "gitlab", "mixed"] + connections: tuple[ObservedSource, ...] = () + repos: tuple[str, ...] + captured_at: datetime + periods: ObservedPeriods + people: tuple[ObservedPerson, ...] + pulls: ObservedPullPeriods + unlinked_branches: tuple[ROIBranchSpend, ...] + unmatched_logins: tuple[str, ...] + + +class ObservedReportResponse(ObservedModel): + report: ObservedReport | None + + +class ObservedIdentityUpdate(BaseModel): + model_config = ConfigDict(extra="forbid") + + email: str + logins: tuple[str, ...] = Field(default=(), max_length=100) + accounts: tuple[ObservedAccount, ...] | None = Field(default=None, max_length=500) + + +class ObservedConnectionIdentities(ObservedSource): + identity_map: Mapping[str, str] + unmatched_logins: tuple[str, ...] + + +class ObservedIdentities(ObservedModel): + gateway_emails: tuple[str, ...] + identity_map: Mapping[str, str] + unmatched_logins: tuple[str, ...] + connections: tuple[ObservedConnectionIdentities, ...] = () + + +class ObservedConnection(ObservedModel): + id: str = "" + source_provider: Literal["github", "gitlab"] + api_url: str + repos: tuple[str, ...] + has_token: bool + update_interval_minutes: float + ready: bool + connection_type: Literal["token", "app"] + + +class ObservedSettings(ObservedConnection): + connections: tuple[ObservedConnection, ...] = () + + +class ObservedSettingsUpdate(BaseModel): + model_config = ConfigDict(extra="forbid") + + connection_id: str | None = Field(default=None, max_length=100) + source_provider: Literal["github", "gitlab"] + api_url: str + token: str | None = None + repos: tuple[str, ...] + update_interval_minutes: float | None = Field(default=None, ge=0, le=43200, allow_inf_nan=False) + + +class ObservedApp(ObservedModel): + configured: bool + can_install: bool = False + api_url: str | None = None + callback_url: str | None = None + + +class ObservedApps(ObservedModel): + github: ObservedApp + gitlab: ObservedApp + + +class ObservedAuthorization(ObservedModel): + url: str diff --git a/tests/integration/README.md b/tests/integration/README.md index e4cb3d96d59..ef418c50759 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -2,6 +2,8 @@ These tests exercise a running gateway, PostgreSQL and Redis with an owned local upstream. CircleCI owns this suite. Tests are grouped by behavior, with no automatic test retries or fallback to paid provider calls +The ROI database contracts in `database/test_roi_observed.py` run in the GitHub Actions `roi-database` Postgres shard and upload coverage on each PR. They own temporary databases and script only the external provider transport. `GITHUB_FILES` in `run.py` assigns these files to GitHub Actions and excludes them from the CircleCI selection + The `cost` group is driven by `cost_tracking_cases.json`, which contains the cost map, literal requests, literal provider responses and expected accounting values. Each case has a name, contract ID, cost-map model, optional deployment overrides, request body, tagged response and exact or recount expectations. Request bodies use `$MODEL` for the registered proxy model, while responses use `$REQUEST_ID` for the per-run scenario ID. To add a case, add a cost-map entry when the model is new, add the request body and exact provider response data, and add hand-computed expected values. The upstream serves each stored response for any path under `/`, while the test-owned cost map is served over loopback through `LITELLM_MODEL_COST_MAP_URL` Use `tests/integration/run.py management`, `accounting`, `database`, `providers`, `extensions`, `mcp`, `sdk` or `cost` to run a selected group. The group to directory mapping is the `GROUPS` literal at the top of `run.py`; a new directory needs a `GROUPS` entry and an `OWNED_DIRECTORIES` entry in `_support/manifest.py`. Set `INTEGRATION_WORKERS` above 1 to run a group under pytest-xdist; the `mcp` job does this in CI, so MCP tests must own their resources per scenario. Set `INTEGRATION_PROXY_URL`, `INTEGRATION_UPSTREAM_URL`, `INTEGRATION_MASTER_KEY` and `DATABASE_URL` to an isolated test deployment. When that deployment runs more than one proxy worker, set `INTEGRATION_PROXY_WORKERS` to the count so a test that writes a model and then calls it waits out the config reload interval, the only cross-worker convergence bound the wire exposes. The runner selects the new domain directories explicitly; the legacy OCI and sandbox selections remain separate @@ -12,7 +14,7 @@ The generated lifecycle models use 20 examples, eight steps, generation and shri Reuse the existing canned provider handlers through `_support/upstream.py`. It rejects internal request fields and exposes actual received requests for independent assertions. Register every created resource for cleanup immediately, keep expected values independent of production calculations, and assert readback plus the runtime effect of a change -The CircleCI workflow starts its own database and Redis, restricts test-phase egress to its owned services and writes JUnit plus an executed-node manifest. Missing setup, failed cleanup or a selected test with neither a passed call nor a skip fail qualification. Skipped nodes are listed under `skipped` in `execution.json`, so the skip reasons double as the open bug list. Existing GitHub Actions jobs do not own these tests +The CircleCI workflow starts its own database and Redis, restricts test-phase egress to its owned services and writes JUnit plus an executed-node manifest. Missing setup, failed cleanup or a selected test with neither a passed call nor a skip fail qualification. Skipped nodes are listed under `skipped` in `execution.json`, so the skip reasons double as the open bug list. GitHub Actions runs only the explicit `GITHUB_FILES` set in `run.py` There is no per-node manifest. A positional argument is a file of the group or a pytest node id inside one (`path::test[param]`), so one cell of a parametrized file can run alone. The runner fails only when pytest fails, when collection errors, or when a selected file collects zero tests. Older tests still carry `@pytest.mark.covers(...)` decorators; the marker stays registered so they collect, but the IDs are not checked against anything and new tests should not use it. The GitHub Actions coverage census reads the `GROUPS` literal in `run.py` and treats every `tests/integration//test_*.py` file in a scheduled group as owned by CircleCI diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 1b39eb81b01..14e31441a95 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -16,6 +16,7 @@ from tests.integration._support.client import Gateway, eventually, gateway_from_ from tests.integration._support.generation import LIFECYCLE_SETTINGS from tests.integration._support.manifest import OWNED_DIRECTORIES from tests.integration._support.routing import RoutingPlugin +from tests.integration.run import GITHUB_FILES COLLECTED: Final = pytest.StashKey[tuple[str, ...]]() REPORTS: Final = pytest.StashKey[list[pytest.TestReport]]() @@ -69,7 +70,10 @@ def pytest_collection_modifyitems(config: pytest.Config, items: list[pytest.Item for item in items if item.path.is_relative_to(root) and item.path.relative_to(root).parts[0] in OWNED_DIRECTORIES ) - if owned and os.environ.get("GITHUB_ACTIONS") == "true": + circleci_only: Final = tuple( + item for item in owned if item.path.relative_to(root.parents[1]).as_posix() not in GITHUB_FILES + ) + if circleci_only and os.environ.get("GITHUB_ACTIONS") == "true": raise pytest.UsageError("Integration contracts are owned by CircleCI") for item in owned: item.add_marker(pytest.mark.integration) diff --git a/tests/integration/database/test_roi_observed.py b/tests/integration/database/test_roi_observed.py new file mode 100644 index 00000000000..f64b09e85c6 --- /dev/null +++ b/tests/integration/database/test_roi_observed.py @@ -0,0 +1,823 @@ +import asyncio +from collections.abc import AsyncIterator +from datetime import date, datetime, timezone +from typing import Final +from urllib.parse import parse_qs, urlsplit + +import httpx +import pytest +import pytest_asyncio +from fastapi import FastAPI, HTTPException +from pydantic import SecretStr + +from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache +from litellm.proxy.roi_calculator.github import SourceError +from litellm.proxy.roi_calculator.oauth import ( + OAuthConfig, + TokenGrant, + begin_authorization, + connected_settings, + consume_state, + exchange_code, + save_grant, +) +from litellm.proxy.roi_calculator.observed_sync import ObservedSyncManager, Progress +from litellm.proxy.roi_calculator.settings import load_settings, load_stored_settings, save_settings +from litellm.proxy.roi_calculator.sync_store import SyncStore +from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.repositories.config_repository import ConfigRepository +from litellm.types.roi_observed import ObservedData, ObservedPeriodData, ObservedWindow +from tests.integration._support.database import scratch_database, write_rows + + +@pytest_asyncio.fixture(loop_scope="function") +async def repository(monkeypatch: pytest.MonkeyPatch) -> AsyncIterator[ConfigRepository]: + with scratch_database() as url: + write_rows( + 'CREATE TABLE "LiteLLM_Config" (param_name TEXT PRIMARY KEY, param_value JSONB NOT NULL, ' + "last_run_at TIMESTAMP NOT NULL DEFAULT NOW(), reload_revision BIGINT NOT NULL DEFAULT 0)", + (), + database_url=url, + ) + monkeypatch.setenv("DATABASE_URL", url) + monkeypatch.delenv("DATABASE_URL_READ_REPLICA", raising=False) + monkeypatch.setenv("LITELLM_SALT_KEY", "test-only-observed-roi-salt-0123456789") + client: Final = PrismaClient(url, ProxyLogging(UserApiKeyCache())) + await client.connect() + try: + yield ConfigRepository(client, use_writer=True) + finally: + await client.disconnect() + + +def _config(provider: str = "github") -> OAuthConfig: + from pydantic import TypeAdapter + + from litellm.proxy.roi_calculator.oauth import Provider + + selected: Final = TypeAdapter(Provider).validate_python(provider) + return OAuthConfig( + selected, + "https://api.github.com" if provider == "github" else "https://gitlab.com/api/v4", + f"https://{provider}.com", + "test-client", + SecretStr("test-client-secret"), + "https://gateway.example.test", + ) + + +def _report() -> ObservedData: + period: Final = ObservedPeriodData( + window=ObservedWindow(start=date(2026, 9, 1), end=date(2026, 9, 28)), pulls=(), issues=(), spend=() + ) + return ObservedData( + source_provider="github", + source_api_url="https://api.github.com", + repos=("org/repo",), + captured_at=datetime.now(timezone.utc), + gateway_emails=(), + current=period, + previous=period, + last_year=period, + ) + + +@pytest.mark.asyncio +async def test_settings_compare_and_swap_rejects_stale_writers(repository: ConfigRepository) -> None: + assert await repository.set_param_if_revision("settings", {"revision": 1, "account": "ari"}, 0) + assert not await repository.set_param_if_revision("settings", {"revision": 1, "account": "bea"}, 0) + writes: Final = await asyncio.gather( + *(repository.set_param_if_revision("settings", {"revision": 2, "account": name}, 1) for name in ("bea", "cam")) + ) + assert sum(writes) == 1 + saved: Final = await repository.get_param("settings") + assert saved is not None and saved.param_value in ( + {"revision": 2, "account": "bea"}, + {"revision": 2, "account": "cam"}, + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider", ("github", "gitlab")) +async def test_authorization_uses_pkce_single_use_state_and_encrypted_credentials( + repository: ConfigRepository, provider: str +) -> None: + config: Final = _config(provider) + url, nonce = await begin_authorization(repository, config) + params: Final = parse_qs(urlsplit(url).query) + assert params["code_challenge_method"] == ["S256"] + assert params["redirect_uri"] == [config.redirect_uri] + state: Final = await consume_state(repository, params["state"][0], nonce, config) + assert state.verifier.get_secret_value() != "**********" + + def respond(request: httpx.Request) -> httpx.Response: + body: Final = parse_qs(request.content.decode()) + assert body["code_verifier"] == [state.verifier.get_secret_value()] + assert body["code"] == ["test-code"] + assert body["client_secret"] == ["test-client-secret"] + assert body["redirect_uri"] == [config.redirect_uri] + return httpx.Response( + 200, json={"access_token": "test-access", "refresh_token": "test-refresh", "expires_in": 3600} + ) + + grant: Final = await exchange_code(config, state, "test-code", httpx.MockTransport(respond)) + await save_grant(repository, config, grant, revision=state.settings_revision) + stored: Final = await load_stored_settings(repository) + assert "test-access" not in stored.model_dump_json() and "test-refresh" not in stored.model_dump_json() + connected: Final = await load_settings(repository) + assert ( + connected.github_token if provider == "github" else connected.gitlab_token + ).get_secret_value() == "test-access" + assert connected.oauth_refresh_token.get_secret_value() == "test-refresh" + with pytest.raises(HTTPException, match="already used"): + await consume_state(repository, params["state"][0], nonce, config) + + +@pytest.mark.asyncio +async def test_authorization_rejects_a_different_browser_and_changed_settings(repository: ConfigRepository) -> None: + config: Final = _config() + url, nonce = await begin_authorization(repository, config) + state: Final = parse_qs(urlsplit(url).query)["state"][0] + with pytest.raises(HTTPException, match="same browser"): + await consume_state(repository, state, "another-browser", config) + new_url, new_nonce = await begin_authorization(repository, config) + verified: Final = await consume_state(repository, parse_qs(urlsplit(new_url).query)["state"][0], new_nonce, config) + await save_grant(repository, config, TokenGrant(access_token=SecretStr("first"))) + with pytest.raises(HTTPException, match="changed during authorization"): + await save_grant( + repository, config, TokenGrant(access_token=SecretStr("stale")), revision=verified.settings_revision + ) + assert (await load_settings(repository)).github_token.get_secret_value() == "first" + assert nonce != new_nonce + + +@pytest.mark.asyncio +async def test_concurrent_app_refresh_rotates_once_and_preserves_account_links( + repository: ConfigRepository, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_ROI_GITHUB_CLIENT_ID", "test-client") + monkeypatch.setenv("LITELLM_ROI_GITHUB_CLIENT_SECRET", "test-client-secret") + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.test") + await save_grant( + repository, + _config(), + TokenGrant(access_token=SecretStr("old-access"), refresh_token=SecretStr("old-refresh"), expires_in=1), + ) + stored: Final = await load_stored_settings(repository) + settings: Final = (await load_settings(repository)).model_copy(update={"identity_map": {"ari": "ari@example.test"}}) + await save_settings(repository, settings, stored.github_token, stored.estimator_key, revision=stored.revision) + requests: Final[asyncio.Queue[httpx.Request]] = asyncio.Queue() + + def respond(request: httpx.Request) -> httpx.Response: + requests.put_nowait(request) + assert parse_qs(request.content.decode())["refresh_token"] == ["old-refresh"] + return httpx.Response( + 200, json={"access_token": "new-access", "refresh_token": "new-refresh", "expires_in": 3600} + ) + + results: Final = await asyncio.gather( + *(connected_settings(repository, httpx.MockTransport(respond)) for _ in range(12)) + ) + assert requests.qsize() == 1 + assert all(result.github_token.get_secret_value() == "new-access" for result in results) + assert all(result.identity_map == {"ari": "ari@example.test"} for result in results) + assert (await load_settings(repository)).oauth_refresh_token.get_secret_value() == "new-refresh" + + +@pytest.mark.asyncio +async def test_refresh_cannot_restore_a_connection_replaced_while_the_provider_responds( + repository: ConfigRepository, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("LITELLM_ROI_GITHUB_CLIENT_ID", "test-client") + monkeypatch.setenv("LITELLM_ROI_GITHUB_CLIENT_SECRET", "test-client-secret") + monkeypatch.setenv("PROXY_BASE_URL", "https://gateway.example.test") + await save_grant( + repository, + _config(), + TokenGrant(access_token=SecretStr("old"), refresh_token=SecretStr("refresh"), expires_in=1), + ) + started: Final = asyncio.Event() + release: Final = asyncio.Event() + + async def respond(request: httpx.Request) -> httpx.Response: + started.set() + await release.wait() + return httpx.Response(200, json={"access_token": "late-token", "expires_in": 3600}) + + pending: Final = asyncio.create_task(connected_settings(repository, httpx.MockTransport(respond))) + await asyncio.wait_for(started.wait(), 2) + await save_grant(repository, _config(), TokenGrant(access_token=SecretStr("replacement"), expires_in=3600)) + release.set() + result: Final = await asyncio.wait_for(pending, 2) + assert result.github_token.get_secret_value() == "replacement" + assert (await load_settings(repository)).github_token.get_secret_value() == "replacement" + + +@pytest.mark.asyncio +async def test_failed_cancelled_and_stale_workers_cannot_replace_the_published_report( + repository: ConfigRepository, +) -> None: + store: Final = SyncStore(repository.prisma_client, "roi_observed") + manager: Final = ObservedSyncManager() + report: Final = _report() + await repository.set_param("roi_observed_report", report.model_dump(mode="json")) + release: Final = asyncio.Event() + + async def build(progress: Progress) -> ObservedData: + await release.wait() + raise SourceError("source failure") + + assert await manager.start(build, store) + assert not await ObservedSyncManager().start(build, store) + await store.cancel() + release.set() + await manager.cancel() + published: Final = await repository.get_param("roi_observed_report") + assert published is not None and ObservedData.model_validate(published.param_value) == report + + async def complete(progress: Progress) -> ObservedData: + return report + + assert await manager.start(complete, store) + + async def finished() -> None: + for _ in range(200): + status: Final = await store.status() + if status is not None and not status.running: + assert status.phase == "complete", status + return + await asyncio.sleep(0.01) + pytest.fail("Observed report did not publish") + + await asyncio.wait_for(finished(), 5) + saved: Final = await repository.get_param("roi_observed_report") + assert saved is not None and ObservedData.model_validate(saved.param_value) == report + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider,install_first", (("github", False), ("github", True), ("gitlab", False))) +async def test_app_callback_round_trip_and_admin_authorization( + repository: ConfigRepository, monkeypatch: pytest.MonkeyPatch, provider: str, install_first: bool +) -> None: + from fastapi import FastAPI + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.management_endpoints.roi_observed_endpoints import ( + get_oauth_repository, + get_observed_transport, + router, + ) + from litellm.proxy.roi_calculator.settings import get_roi_config_repository + + config: Final = _config(provider) + monkeypatch.setenv(f"LITELLM_ROI_{provider.upper()}_CLIENT_ID", config.client_id) + monkeypatch.setenv(f"LITELLM_ROI_{provider.upper()}_CLIENT_SECRET", config.client_secret.get_secret_value()) + monkeypatch.setenv("PROXY_BASE_URL", config.proxy_url) + monkeypatch.setenv("LITELLM_ROI_GITHUB_APP_SLUG", "example-roi") + + def respond(request: httpx.Request) -> httpx.Response: + if request.method == "POST": + assert str(request.url) == config.token_url + if parse_qs(request.content.decode()).get("code") == ["rejected-code"]: + return httpx.Response(400, json={"error": "invalid_grant"}) + return httpx.Response( + 200, + json={ + "access_token": "provider-test-access", + "refresh_token": "provider-test-refresh", + "expires_in": 3600, + }, + ) + assert request.headers["Authorization"] == "Bearer provider-test-access" + assert "PRIVATE-TOKEN" not in request.headers + return httpx.Response(200, json=[]) + + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[get_roi_config_repository] = lambda: repository + app.dependency_overrides[get_oauth_repository] = lambda: repository + app.dependency_overrides[get_observed_transport] = lambda: httpx.MockTransport(respond) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url=config.proxy_url) as client: + for role, expected in ( + (LitellmUserRoles.INTERNAL_USER, 403), + (LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, 403), + (LitellmUserRoles.PROXY_ADMIN, 200), + ): + app.dependency_overrides[user_api_key_auth] = lambda role=role: UserAPIKeyAuth(user_role=role) + response: Final = await client.post(f"/roi-calculator/observed/oauth/{provider}/start") + assert response.status_code == expected, response.text + successful: Final = await client.post( + f"/roi-calculator/observed/oauth/{provider}/start", params={"install": str(install_first).lower()} + ) + assert "HttpOnly" in successful.headers["set-cookie"] and "Secure" in successful.headers["set-cookie"] + initial_state: Final = parse_qs(urlsplit(successful.json()["url"]).query)["state"][0] + installed: Final = ( + await client.get("/roi-calculator/observed/oauth/github/installed", params={"state": initial_state}) + if install_first + else None + ) + if install_first: + assert urlsplit(successful.json()["url"]).path == "/apps/example-roi/installations/new" + assert installed is not None + assert installed.status_code == 303, installed.text + assert urlsplit(installed.headers["location"]).path == "/login/oauth/authorize" + assert parse_qs(urlsplit(installed.headers["location"]).query)["state"][0] != initial_state + state: Final = ( + parse_qs(urlsplit(installed.headers["location"]).query)["state"][0] if installed else initial_state + ) + callback: Final = await client.get( + f"/roi-calculator/observed/oauth/{provider}/callback", params={"state": state, "code": "test-code"} + ) + assert callback.status_code == 303, callback.text + assert callback.headers["location"] == config.proxy_url + f"/ui/roi-calculator/?connected={provider}" + saved: Final = await client.get("/roi-calculator/observed/settings") + assert saved.json()["has_token"] is True and saved.json()["connection_type"] == "app" + assert "provider-test-access" not in saved.text and "provider-test-refresh" not in saved.text + replay: Final = await client.get( + f"/roi-calculator/observed/oauth/{provider}/callback", params={"state": state, "code": "test-code"} + ) + assert replay.status_code == 303 + assert replay.headers["location"] == config.proxy_url + "/ui/roi-calculator/?connection_failed=1" + restart: Final = await client.post(f"/roi-calculator/observed/oauth/{provider}/start") + denied_state: Final = parse_qs(urlsplit(restart.json()["url"]).query)["state"][0] + denied: Final = await client.get( + f"/roi-calculator/observed/oauth/{provider}/callback", + params={"state": denied_state, "error": "access_denied"}, + ) + assert denied.status_code == 303 + assert denied.headers["location"] == config.proxy_url + "/ui/roi-calculator/?connection_cancelled=1" + unchanged: Final = await client.get("/roi-calculator/observed/settings") + assert unchanged.json() == saved.json() + expired_installation: Final = ( + await client.get("/roi-calculator/observed/oauth/github/installed", params={"state": initial_state}) + if install_first + else None + ) + if expired_installation is not None: + assert expired_installation.status_code == 303 + assert expired_installation.headers["location"].endswith("?connection_failed=1") + retry: Final = await client.post(f"/roi-calculator/observed/oauth/{provider}/start") + rejected: Final = await client.get( + f"/roi-calculator/observed/oauth/{provider}/callback", + params={"state": parse_qs(urlsplit(retry.json()["url"]).query)["state"][0], "code": "rejected-code"}, + ) + assert rejected.status_code == 303 + assert rejected.headers["location"] == config.proxy_url + "/ui/roi-calculator/?connection_failed=1" + assert "litellm_roi_oauth" not in client.cookies + assert (await client.get("/roi-calculator/observed/settings")).json() == saved.json() + + +@pytest.mark.asyncio +async def test_identity_api_combines_accounts_and_removes_automatic_links_without_resync( + repository: ConfigRepository, +) -> None: + from fastapi import FastAPI + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.management_endpoints.roi_observed_endpoints import router + from litellm.proxy.roi_calculator.settings import get_roi_config_repository + from litellm.types.roi_observed import ObservedPull + + write_rows('CREATE TABLE "LiteLLM_UserTable" (user_id TEXT PRIMARY KEY, user_email TEXT)', ()) + write_rows( + 'INSERT INTO "LiteLLM_UserTable" VALUES (%s,%s),(%s,%s)', ("ari", "ari@example.test", "bea", "bea@example.test") + ) + settings: Final = (await load_settings(repository)).model_copy(update={"repos": ("org/repo",)}) + await save_settings(repository, settings, "", "") + base: Final = _report() + pulls: Final = tuple( + ObservedPull( + repo="org/repo", + number=index, + title="Change", + url=f"https://github.com/org/repo/pull/{index}", + author=login, + profile_email="ari@example.test" if login == "ari" else "", + merged_at=base.captured_at, + ) + for index, login in enumerate(("ari", "old-ari")) + ) + data: Final = base.model_copy( + update={ + "gateway_emails": ("ari@example.test", "bea@example.test"), + "current": base.current.model_copy(update={"pulls": pulls}), + } + ) + await repository.set_param("roi_observed_report", data.model_dump(mode="json")) + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[get_roi_config_repository] = lambda: repository + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + async with httpx.AsyncClient( + transport=httpx.ASGITransport(app=app), base_url="https://gateway.example.test" + ) as client: + linked: Final = await client.put( + "/roi-calculator/observed/identities", json={"email": "ari@example.test", "logins": ["ari", "old-ari"]} + ) + assert linked.status_code == 200, linked.text + person: Final = linked.json()["report"]["people"][0] + assert person["logins"] == ["ari", "old-ari"] and person["periods"]["current"]["merged_prs"] == 2 + conflict: Final = await client.put( + "/roi-calculator/observed/identities", json={"email": "bea@example.test", "logins": ["old-ari"]} + ) + assert conflict.status_code == 409 + unknown: Final = await client.put( + "/roi-calculator/observed/identities", json={"email": "missing@example.test", "logins": []} + ) + assert unknown.status_code == 422 + removed: Final = await client.put( + "/roi-calculator/observed/identities", json={"email": "ari@example.test", "logins": []} + ) + assert removed.json()["report"]["people"] == [] + assert [login.split(":", 1)[-1] for login in removed.json()["report"]["unmatched_logins"]] == ["ari", "old-ari"] + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY + ) + readonly: Final = await client.get("/roi-calculator/observed/report") + assert readonly.status_code == 200 + denied: Final = await client.put( + "/roi-calculator/observed/identities", json={"email": "ari@example.test", "logins": []} + ) + assert denied.status_code == 403 + + +@pytest.mark.asyncio +async def test_both_app_connections_keep_repositories_credentials_and_account_links( + repository: ConfigRepository, +) -> None: + from litellm.proxy.roi_calculator.settings import connection_id, stored_connections + + await save_grant( + repository, + _config(), + TokenGrant(access_token=SecretStr("github-access"), refresh_token=SecretStr("github-refresh")), + ) + stored: Final = await load_stored_settings(repository) + github: Final = (await load_settings(repository)).model_copy( + update={"repos": ("org/service", "org/docs"), "identity_map": {"ari": "ari@example.test"}} + ) + await save_settings(repository, github, stored.github_token, stored.estimator_key, revision=stored.revision) + await save_grant( + repository, + _config("gitlab"), + TokenGrant(access_token=SecretStr("gitlab-access"), refresh_token=SecretStr("gitlab-refresh")), + ) + gitlab_id: Final = connection_id("gitlab", _config("gitlab").api_url) + github_id: Final = connection_id("github", _config().api_url) + first: Final = await load_settings(repository, selected_id=github_id) + second: Final = await load_settings(repository, selected_id=gitlab_id) + assert first.repos == ("org/service", "org/docs") + assert first.identity_map == {"ari": "ari@example.test"} + assert first.github_token.get_secret_value() == "github-access" + assert first.oauth_refresh_token.get_secret_value() == "github-refresh" + assert second.gitlab_token.get_secret_value() == "gitlab-access" + assert second.oauth_refresh_token.get_secret_value() == "gitlab-refresh" + assert second.identity_map == {} and second.repos == () + await save_grant(repository, _config(), TokenGrant(access_token=SecretStr("github-rotated")), previous=first) + preserved: Final = await load_settings(repository, selected_id=gitlab_id) + assert preserved.model_dump(exclude={"github_token", "github_api_url"}) == second.model_dump( + exclude={"github_token", "github_api_url"} + ) + assert (await load_settings(repository, selected_id=github_id)).github_token.get_secret_value() == "github-rotated" + persisted: Final = await load_stored_settings(repository) + assert len(stored_connections(persisted)) == 2 + assert all( + token not in persisted.model_dump_json() + for token in ("github-access", "github-rotated", "gitlab-access", "github-refresh", "gitlab-refresh") + ) + + +@pytest.mark.asyncio +async def test_cross_provider_account_form_saves_atomically_without_legacy_logins(repository: ConfigRepository) -> None: + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.management_endpoints.roi_observed_endpoints import router + from litellm.proxy.roi_calculator.settings import get_roi_config_repository, stored_connections, write_admin + + write_rows('CREATE TABLE "LiteLLM_UserTable" (user_id TEXT PRIMARY KEY, user_email TEXT)', ()) + write_rows( + 'INSERT INTO "LiteLLM_UserTable" VALUES (%s, %s), (%s, %s)', + ("ari", "ari@example.test", "bea", "bea@example.test"), + ) + await save_grant(repository, _config("github"), TokenGrant(access_token=SecretStr("github-token"))) + await save_grant(repository, _config("gitlab"), TokenGrant(access_token=SecretStr("gitlab-token"))) + connections: Final = stored_connections(await load_stored_settings(repository)) + accounts: Final = tuple({"connection_id": entry.id, "login": "ari"} for entry in connections) + app: Final = FastAPI() + app.include_router(router) + + async def config_repository() -> ConfigRepository: + return repository + + async def admin() -> UserAPIKeyAuth: + return UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + app.dependency_overrides[get_roi_config_repository] = config_repository + app.dependency_overrides[write_admin] = admin + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://gateway.test") as client: + saved: Final = await client.put( + "/roi-calculator/observed/identities", json={"email": "ari@example.test", "accounts": accounts} + ) + assert saved.status_code == 200, saved.text + after: Final = stored_connections(await load_stored_settings(repository)) + assert all(entry.identity_map == {"ari": "ari@example.test"} for entry in after) + conflicting: Final = ({"connection_id": connections[0].id, "login": "bea"}, accounts[1]) + rejected: Final = await client.put( + "/roi-calculator/observed/identities", json={"email": "bea@example.test", "accounts": conflicting} + ) + assert rejected.status_code == 409, rejected.text + assert stored_connections(await load_stored_settings(repository)) == after + + +@pytest.mark.asyncio +async def test_http_sync_combines_providers_retains_period_and_recovers_invalid_reports( + repository: ConfigRepository, +) -> None: + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.management_endpoints.roi_observed_endpoints import ( + get_observed_manager, + get_observed_transport, + router, + ) + from litellm.proxy.roi_calculator.settings import get_roi_config_repository + from litellm.types.roi_observed import ObservedReportResponse, ObservedSettings + + write_rows('CREATE TABLE "LiteLLM_UserTable" (user_id TEXT PRIMARY KEY, user_email TEXT)', ()) + write_rows( + 'CREATE TABLE "LiteLLM_DailyUserSpend" (user_id TEXT, date TEXT, spend DOUBLE PRECISION, api_requests INTEGER)', + (), + ) + write_rows( + 'CREATE TABLE "LiteLLM_SpendLogs" (spend DOUBLE PRECISION, request_tags JSONB, metadata JSONB, "startTime" TIMESTAMP)', + (), + ) + + def respond(request: httpx.Request) -> httpx.Response: + if request.headers.get("Authorization") == "Bearer rejected": + return httpx.Response(401, json={"message": "Bad credentials"}) + if request.url.path == "/graphql": + return httpx.Response( + 200, json={"data": {"search": {"issueCount": 0, "nodes": [], "pageInfo": {"hasNextPage": False}}}} + ) + if request.url.path.startswith("/repos/"): + return httpx.Response(200, json={"full_name": "org/service", "has_issues": True}) + if request.url.path in ("/user/repos", "/api/v4/projects"): + return httpx.Response(200, json=[]) + if request.url.path == "/api/v4/projects/org/service": + return httpx.Response(200, json={"id": 1, "path_with_namespace": "org/service"}) + if request.url.path.endswith(("/merge_requests", "/issues")): + return httpx.Response(200, json=[]) + raise AssertionError(str(request.url)) + + manager: Final = ObservedSyncManager() + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[get_roi_config_repository] = lambda: repository + app.dependency_overrides[get_observed_manager] = lambda: manager + app.dependency_overrides[get_observed_transport] = lambda: httpx.MockTransport(respond) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + async def finish() -> None: + for _ in range(200): + status: Final = await SyncStore(repository.prisma_client, "roi_observed").status() + if status and not status.running: + assert status.phase == "complete", status.error + return + await asyncio.sleep(0.01) + pytest.fail("Sync did not finish") + + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://gateway.test") as client: + assert (await client.get("/roi-calculator/observed/report")).json() == {"report": None} + assert (await client.post("/roi-calculator/observed/sync")).status_code == 409 + for provider, url in (("github", "https://api.github.com"), ("gitlab", "https://gitlab.com/api/v4")): + result: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "source_provider": provider, + "api_url": url, + "token": "test-token", + "repos": ["org/service"], + "update_interval_minutes": 0, + }, + ) + assert result.status_code == 200, result.text + settings: Final = ObservedSettings.model_validate( + (await client.get("/roi-calculator/observed/settings")).json() + ) + assert len(settings.connections) == 2 and all(entry.has_token for entry in settings.connections) + for entry in settings.connections: + repositories: Final = await client.get( + "/roi-calculator/observed/repositories", params={"connection": entry.id} + ) + assert repositories.status_code == 200, repositories.text + rejected: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "source_provider": "github", + "api_url": "https://api.github.com", + "token": "rejected", + "repos": ["org/service"], + }, + ) + assert rejected.status_code == 502, rejected.text + assert ( + ObservedSettings.model_validate((await client.get("/roi-calculator/observed/settings")).json()) == settings + ) + for params in ({"days": 7}, {}): + started: Final = await client.post("/roi-calculator/observed/sync", params=params) + assert started.status_code == 202, started.text + await asyncio.wait_for(finish(), 5) + response: Final = await client.get("/roi-calculator/observed/report") + report: Final = ObservedReportResponse.model_validate(response.json()).report + assert report is not None and report.source_provider == "mixed" + assert len(report.connections) == 2 and report.periods.current.merged_prs == 0 + assert all( + (period.window.end - period.window.start).days == 6 + for period in (report.periods.current, report.periods.previous, report.periods.last_year) + ) + await repository.set_param("roi_observed_report", {"invalid": True}) + assert (await client.get("/roi-calculator/observed/report")).status_code == 500 + recovered: Final = await client.post("/roi-calculator/observed/sync") + assert recovered.status_code == 202, recovered.text + await asyncio.wait_for(finish(), 5) + rebuilt: Final = ObservedReportResponse.model_validate( + (await client.get("/roi-calculator/observed/report")).json() + ).report + assert ( + rebuilt is not None + and (rebuilt.periods.current.window.end - rebuilt.periods.current.window.start).days == 27 + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider", ("github", "gitlab")) +async def test_host_whitespace_preserves_app_credentials_and_account_matches( + repository: ConfigRepository, provider: str +) -> None: + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.management_endpoints.roi_observed_endpoints import get_observed_transport, router + from litellm.proxy.roi_calculator.settings import connection_id, get_roi_config_repository + + config: Final = _config(provider) + granted: Final = await save_grant( + repository, + config, + TokenGrant(access_token=SecretStr("test-access"), refresh_token=SecretStr("test-refresh"), expires_in=3600), + ) + stored: Final = await load_stored_settings(repository) + before: Final = granted.model_copy( + update={"repos": ("org/service",), "identity_map": {"ari": "ari@example.test"}, "ignored_logins": ("bea",)} + ) + await save_settings( + repository, before, stored.github_token, stored.estimator_key, stored.gitlab_token, revision=stored.revision + ) + + def respond(request: httpx.Request) -> httpx.Response: + assert request.headers.get("Authorization") == "Bearer test-access" + if "/repos/" in request.url.path: + return httpx.Response(200, json={"full_name": "org/service"}) + if request.url.path.endswith("/merge_requests"): + return httpx.Response(200, json=[]) + return httpx.Response(200, json={"id": 1, "path_with_namespace": "org/service"}) + + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[get_roi_config_repository] = lambda: repository + app.dependency_overrides[get_observed_transport] = lambda: httpx.MockTransport(respond) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://gateway.test") as client: + response: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "connection_id": connection_id(provider, config.api_url), + "source_provider": provider, + "api_url": f" {config.api_url}/ ", + "repos": ["org/service"], + }, + ) + assert response.status_code == 200, response.text + assert (await load_settings(repository)) == before + + +@pytest.mark.asyncio +async def test_connection_edits_replace_only_the_selected_host_and_keep_the_workspace_schedule( + repository: ConfigRepository, +) -> None: + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + from litellm.proxy.management_endpoints.roi_observed_endpoints import get_observed_transport, router + from litellm.proxy.roi_calculator.settings import get_roi_config_repository, stored_connections + from litellm.types.roi_calculator import ROISettings + from litellm.types.roi_observed import ObservedSettings + + legacy: Final = ROISettings(repos=("org/service",), estimator_model="legacy-model", update_interval_minutes=60) + await save_settings(repository, legacy, "", "") + + def respond(request: httpx.Request) -> httpx.Response: + if "/repos/" in request.url.path: + return httpx.Response(200, json={"full_name": "org/service"}) + if "/projects/" in request.url.path: + return httpx.Response(200, json={"id": 1, "path_with_namespace": "org/service"}) + raise AssertionError(str(request.url)) + + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[get_roi_config_repository] = lambda: repository + app.dependency_overrides[get_observed_transport] = lambda: httpx.MockTransport(respond) + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=app), base_url="http://gateway.test") as client: + assert (await client.get("/roi-calculator/observed/settings")).status_code == 200 + assert (await load_settings(repository)) == legacy + github: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "source_provider": "github", + "api_url": "https://api.github.com", + "token": "github-token", + "repos": ["org/service"], + "update_interval_minutes": 30, + }, + ) + assert github.status_code == 200, github.text + gitlab: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "source_provider": "gitlab", + "api_url": "https://gitlab.com/api/v4", + "token": "gitlab-token", + "repos": ["org/service"], + }, + ) + assert gitlab.status_code == 200, gitlab.text + before: Final = stored_connections(await load_stored_settings(repository)) + edited: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "connection_id": github.json()["id"], + "source_provider": "github", + "api_url": "https://git.example.test/api/v3", + "token": "enterprise-token", + "repos": ["org/service"], + }, + ) + assert edited.status_code == 200, edited.text + settings: Final = ObservedSettings.model_validate(edited.json()) + assert len(settings.connections) == 2 + assert {entry.api_url for entry in settings.connections} == { + "https://git.example.test/api/v3", + "https://gitlab.com/api/v4", + } + assert all(entry.update_interval_minutes == 30 for entry in settings.connections) + after: Final = stored_connections(await load_stored_settings(repository)) + assert next(entry for entry in after if entry.source_provider == "gitlab") == next( + entry for entry in before if entry.source_provider == "gitlab" + ) + assert (await load_settings(repository)).report_mode == "observed" + stale: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "connection_id": github.json()["id"], + "source_provider": "github", + "api_url": "https://api.github.com", + "repos": [], + }, + ) + assert stale.status_code == 404, stale.text + duplicate: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "connection_id": settings.id, + "source_provider": "gitlab", + "api_url": "https://gitlab.com/api/v4", + "repos": [], + }, + ) + assert duplicate.status_code == 409, duplicate.text + assert stored_connections(await load_stored_settings(repository)) == after + manual: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "source_provider": "gitlab", + "api_url": "https://gitlab.com/api/v4", + "repos": ["org/service"], + "update_interval_minutes": 0, + }, + ) + assert manual.status_code == 200, manual.text + reselected: Final = await client.put( + "/roi-calculator/observed/settings", + json={ + "connection_id": settings.id, + "source_provider": "github", + "api_url": "https://git.example.test/api/v3", + "repos": ["org/service"], + }, + ) + assert reselected.status_code == 200, reselected.text + assert all( + entry.update_interval_minutes == 0 + for entry in ObservedSettings.model_validate(reselected.json()).connections + ) diff --git a/tests/integration/run.py b/tests/integration/run.py index c5facec0bd2..aca21ec522c 100644 --- a/tests/integration/run.py +++ b/tests/integration/run.py @@ -23,6 +23,7 @@ GROUPS: Final = MappingProxyType( "security": ("security",), } ) +GITHUB_FILES: Final = frozenset({"tests/integration/database/test_roi_observed.py"}) @dataclass(frozen=True, slots=True) @@ -63,6 +64,7 @@ def main() -> int: str(path.relative_to(root)) for folder in GROUPS[options.group] for path in sorted((root / "tests/integration" / folder).rglob("test_*.py")) + if str(path.relative_to(root)) not in GITHUB_FILES ) if options.list: print("\n".join(group_files)) diff --git a/tests/integration/spend/test_roi_branch_spend.py b/tests/integration/spend/test_roi_branch_spend.py index c90aa0073cd..9efb22c43be 100644 --- a/tests/integration/spend/test_roi_branch_spend.py +++ b/tests/integration/spend/test_roi_branch_spend.py @@ -1,6 +1,7 @@ import json import os import uuid +from collections.abc import Mapping from datetime import date from typing import Final from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit @@ -11,6 +12,8 @@ from prisma import Prisma from psycopg import sql from litellm.proxy.roi_calculator.branch_spend import read_branch_spend +from litellm.types.roi_calculator import ROIBranchSpend +from tests.integration._support.client import Gateway, JsonValue, object_value @pytest.mark.asyncio @@ -46,7 +49,7 @@ async def test_branch_spend_uses_request_tags_once_and_respects_utc_window() -> setup.execute( sql.SQL( 'INSERT INTO {}."LiteLLM_SpendLogs" ("startTime", spend, request_tags) ' - 'VALUES (%s::timestamp, %s, %s::jsonb)' + "VALUES (%s::timestamp, %s, %s::jsonb)" ).format(sql.Identifier(schema)), (timestamp, spend, json.dumps(request_tags)), ) @@ -57,9 +60,9 @@ async def test_branch_spend_uses_request_tags_once_and_respects_utc_window() -> (None, 100, ("litellm-roi-estimator",)), ): setup.execute( - sql.SQL('INSERT INTO {}."LiteLLM_SpendLogs" VALUES (%s::timestamp, %s, %s::jsonb, %s::jsonb)').format( - sql.Identifier(schema) - ), + sql.SQL( + 'INSERT INTO {}."LiteLLM_SpendLogs" VALUES (%s::timestamp, %s, %s::jsonb, %s::jsonb)' + ).format(sql.Identifier(schema)), ( "2026-09-15 00:00:00", spend, @@ -77,3 +80,82 @@ async def test_branch_spend_uses_request_tags_once_and_respects_utc_window() -> assert costs == {"feature/one": (18, 3), "Feature/one": (7, 1), "free": (0, 1)} finally: setup.execute(sql.SQL("DROP SCHEMA {} CASCADE").format(sql.Identifier(schema))) + + +def test_documented_header_and_body_tags_reach_recorded_branch_and_pr_cost(gateway: Gateway) -> None: + import asyncio + from datetime import datetime, timezone + + from litellm.proxy.roi_calculator.branch_spend import attribute_branch_keys + from tests.integration._support.client import eventually + from tests.integration._support.database import read_rows + from tests.integration._support.wire import Reply, Request, wire_server + + marker: Final = uuid.uuid4().hex + repo: Final = f"github.com/integration/{marker}" + branch: Final = "feature/tag-attribution" + tags: Final = [f"repo:{repo}", f"branch:{branch}"] + + def respond(request: Request) -> Reply: + body: Final = object_value(json.loads(request.body)) + assert "tags" not in body and "x-litellm-tags" not in request.headers + return Reply( + body=json.dumps( + { + "id": f"chatcmpl-{uuid.uuid4().hex}", + "object": "chat.completion", + "created": 1, + "model": "owned-model", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 5, "completion_tokens": 3, "total_tokens": 8}, + } + ).encode() + ) + + with wire_server(respond) as upstream, gateway.scenario() as scenario: + model: Final = scenario.model( + api_base=upstream.url + "/v1", input_cost_per_token=0.001, output_cost_per_token=0.002 + ) + examples: Final[tuple[tuple[Mapping[str, JsonValue], Mapping[str, str]], ...]] = ( + ({"metadata": {"tags": tags}}, {}), + ({"tags": tags}, {}), + ({}, {"x-litellm-tags": ", ".join(tags + tags)}), + ) + for payload, headers in examples: + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "tag attribution"}], + **payload, + }, + headers=headers, + ) + assert response.status_code == 200, response.text + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_tags @> %s::jsonb', (json.dumps(tags),) + ), + lambda values: len(values) == 3, + seconds=70, + ) + expected: Final = 3 * (5 * 0.001 + 3 * 0.002) + assert sum(float(row["spend"]) for row in rows) == pytest.approx(expected) + + async def recorded() -> tuple[ROIBranchSpend, ...]: + database: Final = Prisma() + await database.connect() + try: + today: Final = datetime.now(timezone.utc).date() + return await read_branch_spend(database, today, today, (repo,), casefold_repo=True) + finally: + await database.disconnect() + + spending: Final = asyncio.run(recorded()) + costs: Final = attribute_branch_keys(((repo, 1, repo, branch),), spending) + assert costs[(repo, 1)].spend == pytest.approx(expected) + assert costs[(repo, 1)].requests == 3 + assert costs[(repo, 1)].status == "matched" diff --git a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py index 01d40f94d48..9c82dbac3f3 100644 --- a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py @@ -120,6 +120,15 @@ class _ConfigRepository: self.values = MappingProxyType({**self.values, param_name: param_value}) return self.values[param_name] + async def set_param_if_revision(self, param_name: str, param_value: object, revision: int) -> bool: + from litellm.proxy.roi_calculator.settings import StoredROISettings + + stored: Final = StoredROISettings.model_validate(self.values.get(param_name, {})) + if stored.revision != revision: + return False + await self.set_param(param_name, param_value) + return True + def _client( role: LitellmUserRoles, repository: _ConfigRepository, transport: httpx.AsyncBaseTransport | None = None @@ -322,8 +331,14 @@ def test_schedule_rejects_intervals_under_five_minutes(interval: float) -> None: @pytest.mark.parametrize("anchor", ("2026-09-30T12:00:00", "2026-09-30T12:00:00Z", "2026-09-30T14:00:00+02:00")) -def test_schedule_normalizes_legacy_and_offset_timestamps(anchor: str) -> None: - settings: Final = ROISettings(repos=("example/repo",), estimator_model="estimator", update_interval_minutes=60) +@pytest.mark.parametrize("observed", (False, True)) +def test_schedule_normalizes_timestamps_and_respects_report_mode(anchor: str, observed: bool) -> None: + settings: Final = ROISettings( + repos=("example/repo",), + estimator_model="estimator", + update_interval_minutes=60, + report_mode="observed" if observed else "legacy", + ) status: Final = ROISyncStatus( running=False, phase="error", @@ -337,7 +352,8 @@ def test_schedule_normalizes_legacy_and_offset_timestamps(anchor: str) -> None: finished_at=anchor, ) report: Final = sample_report(datetime(2026, 9, 30, tzinfo=timezone.utc)) - assert _next_update(settings, status, report) == datetime(2026, 9, 30, 13, tzinfo=timezone.utc) + expected: Final = None if observed else datetime(2026, 9, 30, 13, tzinfo=timezone.utc) + assert _next_update(settings, status, report) == expected def test_manual_match_recalculates_saved_report_and_removal_restores_cohort() -> None: diff --git a/tests/unit/proxy/roi_calculator/test_github_observed.py b/tests/unit/proxy/roi_calculator/test_github_observed.py new file mode 100644 index 00000000000..8dfaac5d271 --- /dev/null +++ b/tests/unit/proxy/roi_calculator/test_github_observed.py @@ -0,0 +1,148 @@ +import asyncio +import re +from datetime import date, datetime, timedelta, timezone +from typing import Final +from unittest.mock import AsyncMock + +import httpx +import pytest +from pydantic import BaseModel, SecretStr + +from litellm.proxy.roi_calculator.github import SourceError +from litellm.proxy.roi_calculator.github_observed import GitHubObserved +from litellm.types.roi_calculator import ROISettings + + +class _Variables(BaseModel): + q: str + after: str | None + + +class _Query(BaseModel): + variables: _Variables + + +def _node(number: int, merged: datetime) -> dict[str, object]: + return { + "number": number, + "url": f"https://github.com/org/repo/pull/{number}", + "title": "Change", + "createdAt": (merged - timedelta(seconds=16)).isoformat(), + "updatedAt": merged.isoformat(), + "mergedAt": merged.isoformat(), + "author": {"login": "ari", "__typename": "User"}, + } + + +def _page(nodes: tuple[dict[str, object], ...], count: int, cursor: str | None = None) -> httpx.Response: + return httpx.Response( + 200, + json={ + "data": { + "search": { + "issueCount": count, + "nodes": nodes, + "pageInfo": {"hasNextPage": cursor is not None, "endCursor": cursor}, + } + } + }, + ) + + +@pytest.mark.asyncio +async def test_large_history_splits_the_search_limit_without_losing_midnight_or_split_boundaries() -> None: + start: Final = datetime(2026, 9, 1, tzinfo=timezone.utc) + timestamps: Final = tuple(start + timedelta(seconds=index * 60) for index in range(1001)) + + def respond(request: httpx.Request) -> httpx.Response: + assert request.headers["Authorization"] == "Bearer test-only-token" + assert request.url.path == "/graphql" + query: Final = _Query.model_validate_json(request.content).variables + bounds: Final = re.search(r"merged:([^ ]+)\.\.([^ ]+)", query.q) + assert bounds is not None + lower, upper = (datetime.fromisoformat(value.replace("Z", "+00:00")) for value in bounds.groups()) + assert query.q.count("merged:") == 1 + matching: Final = tuple( + _node(index, timestamp) for index, timestamp in enumerate(timestamps) if lower <= timestamp <= upper + ) + offset: Final = int(query.after or 0) + next_cursor: Final = str(offset + 100) if offset + 100 < len(matching) else None + return _page(matching[offset : offset + 100], len(matching), next_cursor) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + source: Final = GitHubObserved(ROISettings(github_token=SecretStr("test-only-token")), client) + pulls: Final = await source.pulls("org/repo", start.date(), start.date()) + assert tuple(pull.number for pull in pulls) == tuple(range(1001)) + assert all( + pull.created_at and (datetime.fromisoformat(pull.merged_at or "") - pull.created_at).total_seconds() == 16 + for pull in pulls + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ("count", "duplicate", "cursor", "partial")) +async def test_incomplete_source_results_fail_instead_of_publishing_understated_counts(failure: str) -> None: + node: Final = _node(1, datetime(2026, 9, 1, tzinfo=timezone.utc)) + + def respond(request: httpx.Request) -> httpx.Response: + if failure == "partial": + return httpx.Response(200, json={"data": None, "errors": [{"message": "permission denied"}]}) + if failure == "duplicate": + return _page((node, node), 2) + if failure == "cursor": + return _page((node,), 2, "repeated") + return _page((node,), 2) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + source: Final = GitHubObserved(ROISettings(), client) + with pytest.raises(SourceError): + await source.pulls("org/repo", date(2026, 9, 1), date(2026, 9, 1)) + + +@pytest.mark.asyncio +async def test_disabled_issue_tracking_is_unknown_instead_of_zero_bugs() -> None: + def respond(request: httpx.Request) -> httpx.Response: + assert request.method == "GET" + return httpx.Response(200, json={"has_issues": False}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + source: Final = GitHubObserved(ROISettings(), client) + assert await source.issues("org/repo", date(2026, 9, 1), date(2026, 9, 1)) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("failure", ("timeout", "unavailable", "rate_limit")) +async def test_read_queries_recover_from_temporary_provider_failures( + failure: str, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(asyncio, "sleep", AsyncMock()) + responses: Final = iter((False, False, True)) + node: Final = _node(1, datetime(2026, 9, 1, tzinfo=timezone.utc)) + + def respond(request: httpx.Request) -> httpx.Response: + if next(responses): + return _page((node,), 1) + if failure == "timeout": + raise httpx.ReadTimeout("scripted timeout", request=request) + return httpx.Response(429 if failure == "rate_limit" else 502) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + source: Final = GitHubObserved(ROISettings(), client) + pulls: Final = await source.pulls("org/repo", date(2026, 9, 1), date(2026, 9, 1)) + assert tuple(pull.number for pull in pulls) == (1,) + + +@pytest.mark.asyncio +async def test_read_retries_stop_after_three_attempts(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(asyncio, "sleep", AsyncMock()) + requests: Final[asyncio.Queue[httpx.Request]] = asyncio.Queue() + + def respond(request: httpx.Request) -> httpx.Response: + requests.put_nowait(request) + return httpx.Response(503) + + async with httpx.AsyncClient(transport=httpx.MockTransport(respond)) as client: + source: Final = GitHubObserved(ROISettings(), client) + with pytest.raises(SourceError, match="HTTP 503"): + await source.pulls("org/repo", date(2026, 9, 1), date(2026, 9, 1)) + assert requests.qsize() == 3 diff --git a/tests/unit/proxy/roi_calculator/test_gitlab.py b/tests/unit/proxy/roi_calculator/test_gitlab.py index 260a1bd9b7e..e3b237a1d41 100644 --- a/tests/unit/proxy/roi_calculator/test_gitlab.py +++ b/tests/unit/proxy/roi_calculator/test_gitlab.py @@ -279,3 +279,28 @@ async def test_gitlab_retries_transient_errors_and_checks_merge_request_access() assert next(statuses, None) is None finally: await source.close() + + +@pytest.mark.asyncio +async def test_observed_issues_preserve_last_second_boundaries_and_disabled_tracking() -> None: + def respond(request: httpx.Request) -> httpx.Response: + if request.url.path.endswith("/projects/org/disabled"): + return httpx.Response(200, json={"id": 2, "path_with_namespace": "org/disabled", "issues_enabled": False}) + if request.url.path.endswith("/projects/org/repo"): + return httpx.Response(200, json={"id": 1, "path_with_namespace": "org/repo"}) + assert request.url.params["created_before"] == "2026-10-01T00:00:00Z" + return httpx.Response( + 200, + json=[ + {"iid": 1, "created_at": "2026-09-30T23:59:59.999Z", "labels": ["type::bug"]}, + {"iid": 2, "created_at": "2026-10-01T00:00:00Z", "labels": ["bug"]}, + ], + ) + + source: Final = GitLab(ROISettings(source_provider="gitlab"), httpx.MockTransport(respond)) + try: + issues: Final = await source.issues("org/repo", date(2026, 9, 1), date(2026, 9, 30)) + assert issues is not None and tuple(issue.number for issue in issues) == (1,) + assert await source.issues("org/disabled", date(2026, 9, 1), date(2026, 9, 30)) is None + finally: + await source.close() diff --git a/tests/unit/proxy/roi_calculator/test_oauth.py b/tests/unit/proxy/roi_calculator/test_oauth.py new file mode 100644 index 00000000000..a8b8716fc4c --- /dev/null +++ b/tests/unit/proxy/roi_calculator/test_oauth.py @@ -0,0 +1,25 @@ +from dataclasses import replace +from typing import Final + +from pydantic import SecretStr + +from litellm.proxy.roi_calculator.oauth import OAuthConfig + + +def test_app_urls_support_enterprise_and_a_gateway_path_prefix() -> None: + cloud: Final = OAuthConfig( + "github", + "https://api.github.com", + "https://github.com", + "test-client", + SecretStr("test-secret"), + "https://gateway.example.test/proxy", + "test-app", + ) + enterprise: Final = replace(cloud, api_url="https://git.example.test/api/v3", base_url="https://git.example.test") + assert cloud.installation_url == cloud.base_url + "/apps/test-app/installations/new" + assert enterprise.installation_url == enterprise.base_url + "/github-apps/test-app/installations/new" + assert enterprise.cookie_path == "/proxy/roi-calculator/observed/oauth" + assert enterprise.redirect_uri.startswith(enterprise.proxy_url + "/roi-calculator/") + assert replace(cloud, provider="gitlab").installation_url is None + assert replace(cloud, app_slug="").installation_url is None diff --git a/tests/unit/proxy/roi_calculator/test_observed_analytics.py b/tests/unit/proxy/roi_calculator/test_observed_analytics.py new file mode 100644 index 00000000000..a8416f985da --- /dev/null +++ b/tests/unit/proxy/roi_calculator/test_observed_analytics.py @@ -0,0 +1,176 @@ +from datetime import date, datetime, timedelta, timezone +from typing import Final + +import pytest + +from litellm.proxy.roi_calculator.observed_analytics import ( + declared_requester, + merge_hours, + reporting_windows, + summarize_observed, +) +from litellm.types.roi_observed import ObservedData, ObservedIssue, ObservedPeriodData, ObservedPull, ObservedWindow + +_NOW: Final = datetime(2026, 10, 3, tzinfo=timezone.utc) +_WINDOW: Final = ObservedWindow(start=date(2026, 9, 5), end=date(2026, 10, 2)) + + +def _pull(login: str, number: int = 1, repo: str = "org/service", **fields: object) -> ObservedPull: + return ObservedPull.model_validate( + { + "repo": repo, + "number": number, + "title": "Ship change", + "url": f"https://github.com/{repo}/pull/{number}", + "author": login, + "created_at": "2026-09-10T00:00:00Z", + "merged_at": "2026-09-10T00:01:19Z", + **fields, + } + ) + + +def _data(current: ObservedPeriodData, previous: ObservedPeriodData | None = None) -> ObservedData: + empty: Final = ObservedPeriodData(window=_WINDOW, pulls=(), issues=(), spend=()) + return ObservedData( + source_provider="github", + source_api_url="https://api.github.com", + repos=("org/service",), + captured_at=_NOW, + gateway_emails=("ari@example.test", "bea@example.test"), + current=current, + previous=previous or empty, + last_year=empty, + ) + + +def test_multiple_accounts_share_one_cost_denominator_and_pr_numbers_are_scoped_to_repositories() -> None: + direct: Final = _pull("ari", profile_email="ari@example.test") + alternate: Final = _pull("old-ari", repo="org/other") + agent: Final = _pull("devin-ai[bot]", 3, agent=True, requester="old-ari") + unowned: Final = _pull("devin-ai[bot]", 4, agent=True) + period: Final = ObservedPeriodData( + window=_WINDOW, + pulls=(direct, alternate, agent, unowned), + issues=(), + spend=({"email": "ari@example.test", "spend": 90.0, "date": "2026-09-10", "user_id": "ari", "requests": 1},), + ) + report: Final = summarize_observed(_data(period), {"old-ari": "ari@example.test"}) + assert len(report.people) == 1 + person: Final = report.people[0] + assert person.logins == ("ari", "old-ari") + assert person.periods.current.pr_urls == (direct.url, alternate.url, agent.url) + assert (person.periods.current.direct_authored, person.periods.current.declared_agent_owned) == (2, 1) + assert person.periods.current.recorded_spend_per_attributed_pr == 30.0 + assert person.periods.current.prs_per_week == 0.75 + assert (report.periods.current.merged_prs, report.periods.current.matched_internal_prs) == (4, 3) + assert report.periods.current.agents_without_requester == 1 + + +def test_missing_spend_stays_unknown_and_a_recorded_zero_stays_zero() -> None: + period: Final = ObservedPeriodData( + window=_WINDOW, + pulls=(_pull("ari"), _pull("bea", 2)), + issues=None, + spend=({"email": "bea@example.test", "spend": 0.0, "date": "2026-09-10", "user_id": "bea", "requests": 1},), + ) + report: Final = summarize_observed(_data(period), {"ari": "ari@example.test", "bea": "bea@example.test"}) + assert tuple(person.periods.current.recorded_spend_per_attributed_pr for person in report.people) == (None, 0.0) + assert tuple(person.periods.current.spend_observation for person in report.people) == ( + "no_records", + "records_present", + ) + assert report.periods.current.new_bug_labeled_issues is None + assert report.periods.previous.new_bug_labeled_issues == 0 + assert report.people[0].periods.previous.recorded_spend_per_attributed_pr is None + + +def test_manual_links_override_automatic_matches_and_removal_suppresses_rematching() -> None: + period: Final = ObservedPeriodData( + window=_WINDOW, pulls=(_pull("ari", profile_email="ari@example.test"),), issues=(), spend=() + ) + data: Final = _data(period) + assert summarize_observed(data, {}).people[0].email == "ari@example.test" + assert summarize_observed(data, {"ari": "bea@example.test"}).people[0].email == "bea@example.test" + removed: Final = summarize_observed(data, {}, ("ari",)) + assert removed.people == () + assert removed.unmatched_logins == ("ari",) + assert summarize_observed(data, {"ari": "bea@example.test"}, ("ari",)).people[0].email == "bea@example.test" + + +def test_conflicting_public_emails_do_not_silently_choose_an_owner() -> None: + current: Final = ObservedPeriodData( + window=_WINDOW, pulls=(_pull("ari", profile_email="ari@example.test"),), issues=(), spend=() + ) + previous: Final = current.model_copy(update={"pulls": (_pull("ari", profile_email="bea@example.test"),)}) + report: Final = summarize_observed(_data(current, previous), {}) + assert report.people == () + assert report.unmatched_logins == ("ari",) + + +def test_quality_counts_labelled_issues_once_and_does_not_infer_bugs_from_pr_titles() -> None: + period: Final = ObservedPeriodData( + window=_WINDOW, + pulls=(_pull("ari", title="fix: critical bug"), _pull("ari", 2, title='Revert "change"')), + issues=tuple( + ObservedIssue(repo="org/service", number=index, created_at=_NOW, labels=labels) + for index, labels in enumerate( + ( + ("BUG", "kind:bug"), + ("type::bug", "type::regression"), + ("debug",), + ) + ) + ), + spend=(), + ) + report: Final = summarize_observed(_data(period), {}) + assert ( + report.periods.current.new_bug_labeled_issues, + report.periods.current.new_regression_labeled_issues, + report.periods.current.explicitly_titled_revert_prs, + ) == (2, 1, 1) + + +@pytest.mark.parametrize( + "created,merged,expected", + ( + (None, "2026-09-10T00:00:16Z", None), + ("2026-09-10T00:00:00Z", "2026-09-10T00:00:16Z", 16 / 3600), + ("2026-09-10T00:00:00Z", "2026-09-10T00:00:00Z", 0), + ("2026-09-10T00:00:01Z", "2026-09-10T00:00:00Z", None), + ("2026-09-10T00:00:00", "2026-09-10T00:00:16Z", None), + ), +) +def test_merge_duration_preserves_seconds_and_rejects_invalid_intervals( + created: str | None, merged: str, expected: float | None +) -> None: + assert merge_hours(_pull("ari", created_at=created, merged_at=merged)) == expected + + +@pytest.mark.parametrize( + "now", (datetime(2024, 3, 1, tzinfo=timezone.utc), _NOW, _NOW.replace(tzinfo=timezone(timedelta(hours=14)))) +) +@pytest.mark.parametrize("days", (1, 7, 28, 90, 366)) +def test_reporting_windows_have_equal_lengths_and_exclude_today(now: datetime, days: int) -> None: + current, previous, yearly = reporting_windows(now, days) + assert all((window.end - window.start).days + 1 == days for window in (current, previous, yearly)) + assert current.end == now.astimezone(timezone.utc).date() - timedelta(days=1) + assert previous.end == current.start - timedelta(days=1) + assert yearly.end.year == current.end.year - 1 + assert yearly.end.month == current.end.month + + +@pytest.mark.parametrize( + "author,body,expected", + ( + ("devin-ai[bot]", "Requested by: @Ari", "ari"), + ("devin-ai-integration", "Requested by: @Ari", "ari"), + ("devin-ai", "Requested by: @Ari", "ari"), + ("human", "Requested by: @ari", ""), + ("devin-ai[bot]", "Requested by: @ari\nRequested by: @bea", ""), + ("devin-ai[bot]", "Mentions @ari", ""), + ), +) +def test_agent_ownership_requires_one_explicit_requester(author: str, body: str, expected: str) -> None: + assert declared_requester(author, body) == expected diff --git a/tests/unit/proxy/roi_calculator/test_observed_sync.py b/tests/unit/proxy/roi_calculator/test_observed_sync.py new file mode 100644 index 00000000000..55ffe401286 --- /dev/null +++ b/tests/unit/proxy/roi_calculator/test_observed_sync.py @@ -0,0 +1,140 @@ +from datetime import date, datetime, timezone +from typing import Final, Literal + +import httpx +import pytest +from pydantic import BaseModel, SecretStr + +from litellm.proxy.roi_calculator.observed_analytics import summarize_observed +from litellm.proxy.roi_calculator.observed_sync import collect_observed +from litellm.proxy.roi_calculator.source import repository_tag +from litellm.types.roi_calculator import ROIBranchSpend, ROISettings, ROISpendRecord + + +class _Variables(BaseModel): + q: str + + +class _Query(BaseModel): + variables: _Variables + + +@pytest.mark.asyncio +@pytest.mark.parametrize("provider", ("github", "gitlab")) +@pytest.mark.parametrize("days", (7, 28, 90)) +@pytest.mark.parametrize("include_disabled_repo", (False, True)) +async def test_live_provider_metadata_reaches_people_quality_durations_and_branch_spend_without_an_estimator( + provider: Literal["github", "gitlab"], + days: int, + include_disabled_repo: bool, +) -> None: + settings: Final = ROISettings( + source_provider=provider, + repos=("org/repo", "org/disabled") if include_disabled_repo else ("org/repo",), + github_token=SecretStr("source-test-token"), + gitlab_token=SecretStr("source-test-token"), + identity_map={"old-ari": "ari@example.test"}, + ) + tag: Final = repository_tag(settings, "org/repo") + + async def spend(start: date, end: date) -> tuple[ROISpendRecord, ...]: + return ({"date": str(start), "user_id": "ari", "email": "ari@example.test", "spend": 30.0, "requests": 10},) + + async def users() -> frozenset[str]: + return frozenset(("ari@example.test",)) + + async def branches(start: date, end: date, repos: tuple[str, ...]) -> tuple[ROIBranchSpend, ...]: + assert repos == tuple(sorted(repository_tag(settings, repo) for repo in settings.repos)) + return (ROIBranchSpend(repo=tag, branch="fix/parser", spend=5.0, requests=2),) + + def respond(request: httpx.Request) -> httpx.Response: + path: Final = request.url.path + if path == "/graphql": + query: Final = _Query.model_validate_json(request.content).variables.q + if "repo:org/disabled " in query: + return httpx.Response( + 200, json={"data": {"search": {"issueCount": 0, "nodes": [], "pageInfo": {"hasNextPage": False}}}} + ) + kind: Final = "pull" if "is:pr" in query else "issue" + start: Final = query.split("merged:" if kind == "pull" else "created:")[1][:10] + node: Final = { + "number": 1, + "url": "https://github.com/org/repo/pull/1", + "title": "Change", + "createdAt": f"{start}T12:00:00Z", + "updatedAt": f"{start}T12:00:16Z", + "mergedAt": f"{start}T12:00:16Z", + "author": {"login": "devin-ai-integration", "__typename": "Bot"}, + "body": "Requested by: @old-ari", + "headRefName": "fix/parser", + "headRepository": {"nameWithOwner": "org/repo"}, + "labels": {"nodes": [{"name": "bug"}], "pageInfo": {"hasNextPage": False}}, + } + return httpx.Response( + 200, json={"data": {"search": {"issueCount": 1, "nodes": [node], "pageInfo": {"hasNextPage": False}}}} + ) + if path == "/repos/org/repo": + return httpx.Response(200, json={"has_issues": True}) + if path == "/repos/org/disabled": + return httpx.Response(200, json={"has_issues": False}) + if path.endswith("/projects/org/disabled"): + return httpx.Response(200, json={"id": 2, "path_with_namespace": "org/disabled", "issues_enabled": False}) + if path.endswith("/projects/2/merge_requests"): + return httpx.Response(200, json=[]) + if path.endswith("/projects/org/repo"): + return httpx.Response(200, json={"id": 1, "path_with_namespace": "org/repo"}) + if path.endswith("/merge_requests"): + start: Final = request.url.params["merged_after"][:10] + return httpx.Response( + 200, + json=[ + { + "iid": 1, + "web_url": "https://gitlab.com/org/repo/-/merge_requests/1", + "title": "Change", + "author": {"username": "old-ari"}, + "created_at": f"{start}T12:00:00Z", + "updated_at": f"{start}T12:00:16Z", + "merged_at": f"{start}T12:00:16Z", + "source_branch": "fix/parser", + "source_project_id": 1, + } + ], + ) + if path.endswith("/issues"): + start: Final = request.url.params["created_after"][:10] + return httpx.Response(200, json=[{"iid": 1, "created_at": f"{start}T12:00:00Z", "labels": ["bug"]}]) + raise AssertionError(f"Unexpected API request: {request.method} {path}") + + data: Final = await collect_observed( + settings, + spend, + users, + branches, + datetime(2026, 10, 3, tzinfo=timezone.utc), + lambda stage, done, total: None, + httpx.MockTransport(respond), + days=days, + ) + assert all( + (period.window.end - period.window.start).days + 1 == days + for period in (data.current, data.previous, data.last_year) + ) + report: Final = summarize_observed(data, settings.identity_map) + person: Final = report.people[0].periods.current + assert (person.merged_prs, person.gateway_recorded_spend, person.recorded_spend_per_attributed_pr) == ( + 1, + 30.0, + 30.0, + ) + assert person.median_merge_hours == 16 / 3600 + assert person.declared_agent_owned == (1 if provider == "github" else 0) + assert tuple( + window.merged_prs for window in (report.periods.current, report.periods.previous, report.periods.last_year) + ) == (1, 1, 1) + assert tuple( + period.new_bug_labeled_issues + for period in (report.periods.current, report.periods.previous, report.periods.last_year) + ) == (1, 1, 1) + assert report.pulls.current[0].branch_cost.spend == 5.0 + assert report.unlinked_branches == () diff --git a/tests/unit/proxy/roi_calculator/test_observed_workspace.py b/tests/unit/proxy/roi_calculator/test_observed_workspace.py new file mode 100644 index 00000000000..42ea77caf54 --- /dev/null +++ b/tests/unit/proxy/roi_calculator/test_observed_workspace.py @@ -0,0 +1,136 @@ +from datetime import date, datetime, timezone +from typing import Final + +import pytest + +from litellm.proxy.roi_calculator.github import SourceError +from litellm.proxy.roi_calculator.observed_workspace import ( + combine_observed, + scoped_data, + source_details, + summarize_workspace, +) +from litellm.proxy.roi_calculator.settings import StoredConnection +from litellm.types.roi_calculator import ROIBranchSpend, ROISettings +from litellm.types.roi_observed import ObservedData, ObservedIssue, ObservedPeriodData, ObservedPull, ObservedWindow + + +def _source(settings: ROISettings, issues: tuple[ObservedIssue, ...] | None = ()) -> ObservedData: + host: Final = "github.com" if settings.source_provider == "github" else "gitlab.com" + period: Final = ObservedPeriodData( + window=ObservedWindow(start=date(2026, 9, 1), end=date(2026, 9, 28)), + pulls=tuple( + ObservedPull( + repo=repo, + number=1, + title="Change", + url=f"https://{host}/{repo}/pull/1", + author="ari", + created_at=datetime(2026, 9, 10, 0, 0, 0, tzinfo=timezone.utc), + merged_at=datetime(2026, 9, 10, 0, 0, 30, tzinfo=timezone.utc), + source_repo=f"{host}/{repo}", + source_branch="feature/one", + ) + for repo in settings.repos + ), + issues=issues, + spend=({"date": "2026-09-10", "user_id": "ari", "email": "ari@example.test", "spend": 60.0, "requests": 3},), + branch_spend=tuple( + ROIBranchSpend(repo=f"{host}/{repo}", branch="feature/one", spend=2, requests=1) for repo in settings.repos + ), + ) + data: Final = ObservedData( + source_provider=settings.source_provider, + source_api_url=settings.source_api_url, + repos=settings.repos, + captured_at=datetime(2026, 9, 29, tzinfo=timezone.utc), + gateway_emails=("ari@example.test", "bea@example.test"), + current=period, + previous=period, + last_year=period, + ) + return scoped_data(data, source_details(settings)) + + +@pytest.mark.parametrize("same_person", (True, False)) +def test_multiple_repos_and_providers_scope_usernames_and_count_spend_once(same_person: bool) -> None: + github: Final = ROISettings(repos=("org/service", "org/docs")) + gitlab: Final = ROISettings(source_provider="gitlab", repos=("org/service",)) + combined: Final = combine_observed( + (_source(github), _source(gitlab)), ("github.com/org/service", "github.com/org/docs", "gitlab.com/org/service") + ) + report: Final = summarize_workspace( + combined, + ( + StoredConnection( + source_provider="github", api_url=github.source_api_url, identity_map={"ari": "ari@example.test"} + ), + StoredConnection( + source_provider="gitlab", + api_url=gitlab.source_api_url, + identity_map={"ari": "ari@example.test" if same_person else "bea@example.test"}, + ), + ), + ) + person: Final = next(person for person in report.people if person.email == "ari@example.test") + assert report.source_provider == "mixed" + assert report.periods.current.merged_prs == 3 + assert report.periods.current.matched_users_recorded_spend == 60 + assert person.periods.current.merged_prs == (3 if same_person else 2) + assert person.periods.current.gateway_recorded_spend == 60 + assert person.periods.current.recorded_spend_per_attributed_pr == (20 if same_person else 30) + assert len(person.accounts) == (2 if same_person else 1) + assert all(pull.branch_cost.spend == 2 for pull in report.pulls.current) + assert report.unlinked_branches == () + + +def test_empty_repository_is_a_successful_zero_activity_report() -> None: + settings: Final = ROISettings(repos=()) + data: Final = combine_observed((_source(settings),), ("org/empty",)) + report: Final = summarize_workspace(data, ()) + assert report.periods.current.merged_prs == 0 + assert report.periods.current.new_bug_labeled_issues == 0 + assert report.periods.current.median_merge_hours is None + assert report.people == () and report.pulls.current == () + + +@pytest.mark.parametrize( + ("issues", "expected"), + ( + (None, None), + ((), 0), + ( + ( + ObservedIssue( + repo="org/service", + number=1, + created_at=datetime(2026, 9, 10, tzinfo=timezone.utc), + labels=("bug", "regression"), + ), + ), + 1, + ), + ), +) +def test_disabled_tracking_does_not_hide_other_connections_quality_counts( + issues: tuple[ObservedIssue, ...] | None, expected: int | None +) -> None: + github: Final = _source(ROISettings(repos=("org/docs",)), issues=None) + gitlab: Final = _source(ROISettings(source_provider="gitlab", repos=("org/service",)), issues=issues) + report: Final = summarize_workspace(combine_observed((github, gitlab), ("org/docs", "org/service")), ()) + assert tuple( + (period.new_bug_labeled_issues, period.new_regression_labeled_issues) + for period in (report.periods.current, report.periods.previous, report.periods.last_year) + ) == ((expected, expected),) * 3 + assert report.periods.current.merged_prs == 2 + + +def test_duplicate_connection_cannot_double_count_a_merged_change() -> None: + data: Final = _source(ROISettings(repos=("org/service",))) + with pytest.raises(SourceError, match="more than one connection"): + combine_observed((data, data), data.repos) + + +def test_github_repository_selection_deduplicates_case_variants() -> None: + settings: Final = ROISettings(repos=("Org/Service", "org/service", "org/docs", "org/docs.git")) + assert settings.repos == ("Org/Service", "org/docs") diff --git a/tests/unit/test_assert_ci_coverage.py b/tests/unit/test_assert_ci_coverage.py index cc25627c651..5f3903a5feb 100644 --- a/tests/unit/test_assert_ci_coverage.py +++ b/tests/unit/test_assert_ci_coverage.py @@ -73,6 +73,19 @@ def test_integration_groups_require_exclusive_scheduled_circleci_owner(tmp_path: assert [(finding.subject, finding.detail) for finding in findings] == [ (test_path, "integration contract is also selected by GitHub Actions") ] + github_path: Final = "tests/integration/management/test_github_contract.py" + (tmp_path / github_path).write_text("def test_contract(): pass\n") + runner: Final = tmp_path / "tests/integration/run.py" + runner.write_text(runner.read_text() + f"GITHUB_FILES: Final = frozenset({{{github_path!r}}})\n") + workflow.write_text(yaml.safe_dump({"jobs": {"tests": {"steps": [{"run": f"pytest {github_path}"}]}}})) + github_owned, github_findings = coverage._integration_ownership(tmp_path) + assert github_owned == frozenset({test_path, github_path}) + assert github_findings == () + workflow.write_text(yaml.safe_dump({"jobs": {}})) + _, missing_invocation = coverage._integration_ownership(tmp_path) + assert [(finding.subject, finding.detail) for finding in missing_invocation] == [ + (github_path, "GitHub-owned integration contract has no invoking workflow") + ] def test_an_ancestor_directory_covers_a_file_but_does_not_name_it(): diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.integration.test.tsx new file mode 100644 index 00000000000..bc242ebe14a --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.integration.test.tsx @@ -0,0 +1,88 @@ +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { afterEach, expect, it, vi } from "vitest"; +import ObservedAccounts from "./ObservedAccounts"; + +afterEach(() => vi.unstubAllGlobals()); +it("links several usernames on both providers to one email in a single save", async () => { + const writes: unknown[] = []; + const saved = vi.fn(); + vi.stubGlobal( + "fetch", + vi.fn(async (_input: string, init: RequestInit) => { + if (init.method === "PUT") { + writes.push(JSON.parse(String(init.body))); + return Response.json({ report: null }); + } + const identities = { + gateway_emails: ["ari@example.test"], + identity_map: { old: "ari@example.test" }, + unmatched_logins: ["new"], + connections: [ + { + id: "github-id", + source_provider: "github", + api_url: "https://api.github.com", + identity_map: { old: "ari@example.test" }, + unmatched_logins: ["new"], + }, + { + id: "gitlab-id", + source_provider: "gitlab", + api_url: "https://gitlab.com/api/v4", + identity_map: {}, + unmatched_logins: ["new"], + }, + ], + }; + return Response.json(identities); + }), + ); + const user = userEvent.setup(); + render(); + await screen.findByLabelText(/GitHub usernames/); + fireEvent.change(screen.getByLabelText("Internal email"), { target: { value: "ari@example.test" } }); + expect(screen.getByLabelText(/GitHub usernames/)).toHaveValue("old"); + fireEvent.change(screen.getByLabelText(/GitHub usernames/), { target: { value: "@Old, new, NEW" } }); + fireEvent.change(screen.getByLabelText(/GitLab usernames/), { target: { value: "new" } }); + await user.click(screen.getByRole("button", { name: "Save accounts" })); + await waitFor(() => expect(saved).toHaveBeenCalledOnce()); + expect(writes).toEqual([ + { + email: "ari@example.test", + accounts: [ + { connection_id: "github-id", login: "old" }, + { connection_id: "github-id", login: "new" }, + { connection_id: "gitlab-id", login: "new" }, + ], + }, + ]); +}); +it("keeps a conflicting link editable", async () => { + const saved = vi.fn(); + vi.stubGlobal( + "fetch", + vi.fn(async (_input: string, init: RequestInit) => + init.method === "PUT" + ? Response.json({ detail: "An account is already linked to another email. Unlink it first" }, { status: 409 }) + : Response.json({ gateway_emails: ["ari@example.test"], identity_map: {}, unmatched_logins: [] }), + ), + ); + const user = userEvent.setup(); + render( + , + ); + fireEvent.change(screen.getByLabelText("Source usernames"), { target: { value: "old, new" } }); + const button = screen.getByRole("button", { name: "Save accounts" }); + await waitFor(() => expect(button).toBeEnabled()); + await user.click(button); + expect(await screen.findByRole("alert")).toHaveTextContent("An account is already linked to another email"); + expect(screen.getByLabelText("Source usernames")).toHaveValue("old, new"); + expect(saved).not.toHaveBeenCalled(); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.tsx new file mode 100644 index 00000000000..ab6c4cf38df --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedAccounts.tsx @@ -0,0 +1,209 @@ +"use client"; + +import { useEffect, useState } from "react"; +import { z } from "zod"; +import { apiClient } from "@/components/networking"; +import { extractProxyErrorMessage } from "@/lib/http/client"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogDescription } from "@/components/ui/dialog"; +import { accountLogins, type ObservedPerson } from "./observedData"; + +const connectionIdentityFields = { + id: z.string(), + source_provider: z.enum(["github", "gitlab"]), + api_url: z.string(), + identity_map: z.record(z.string(), z.string()), + unmatched_logins: z.array(z.string()), +}; +const identitiesFields = { + gateway_emails: z.array(z.string()), + identity_map: z.record(z.string(), z.string()), + unmatched_logins: z.array(z.string()), + connections: z.array(z.object(connectionIdentityFields)).optional(), +}; +const identitiesSchema = z.object(identitiesFields); + +function matches(identities: z.infer, email: string, people: ObservedPerson[]) { + const person = people.find((entry) => entry.email === email); + return Object.fromEntries( + (identities.connections ?? []).map((entry) => [ + entry.id, + [ + ...new Set([ + ...Object.entries(entry.identity_map) + .filter(([, address]) => address === email) + .map(([login]) => login), + ...(person?.accounts ?? []) + .filter((account) => account.connection_id === entry.id) + .map((account) => account.login), + ]), + ].join(", "), + ]), + ); +} + +export default function ObservedAccounts({ + accessToken, + people, + initialEmail = "", + onClose, + onSaved, +}: { + accessToken: string; + people: ObservedPerson[]; + initialEmail?: string; + onClose: () => void; + onSaved: () => void; +}) { + const [identities, setIdentities] = useState | null>(null); + const [email, setEmail] = useState(initialEmail); + const [logins, setLogins] = useState(people.find((person) => person.email === initialEmail)?.logins.join(", ") ?? ""); + const [linked, setLinked] = useState>({}); + const [saving, setSaving] = useState(false); + const [error, setError] = useState(""); + useEffect(() => { + const controller = new AbortController(); + apiClient + .get("/roi-calculator/observed/identities", { accessToken, signal: controller.signal }) + .then((data) => { + if (!controller.signal.aborted) { + const parsed = identitiesSchema.parse(data); + setIdentities(parsed); + setLinked(matches(parsed, initialEmail, people)); + } + }) + .catch((reason: unknown) => { + if (!controller.signal.aborted) setError(extractProxyErrorMessage(reason)); + }); + return () => controller.abort(); + }, [accessToken, initialEmail, people]); + function selectEmail(value: string) { + setEmail(value); + if (identities) setLinked(matches(identities, value, people)); + const automatic = people.find((person) => person.email === value)?.logins ?? []; + const manual = Object.entries(identities?.identity_map ?? {}) + .filter(([, address]) => address === value) + .map(([login]) => login); + setLogins([...new Set([...automatic, ...manual])].join(", ")); + } + async function save() { + setSaving(true); + setError(""); + try { + await apiClient.put("/roi-calculator/observed/identities", { + accessToken, + body: { + email: email.trim().toLowerCase(), + ...(identities?.connections?.length + ? { + accounts: identities.connections.flatMap((entry) => + accountLogins(linked[entry.id] ?? "").map((login) => ({ connection_id: entry.id, login })), + ), + } + : { logins: accountLogins(logins) }), + }, + }); + onSaved(); + onClose(); + } catch (reason) { + setError(extractProxyErrorMessage(reason)); + } finally { + setSaving(false); + } + } + return ( + { + if (!open) onClose(); + }} + > + + + Link accounts + Match one internal user to all their source accounts + +
+
+ + selectEmail(event.target.value)} + /> + + {identities?.gateway_emails.map((address) => +
+ {identities?.connections?.length ? ( + identities.connections.map((entry) => ( +
+ + setLinked({ ...linked, [entry.id]: event.target.value })} + placeholder="current-account, old-account" + /> + {entry.unmatched_logins.length > 0 && ( +
+ {entry.unmatched_logins.length} unmatched accounts +
+ {entry.unmatched_logins.map((login) => ( + + ))} +
+
+ )} +
+ )) + ) : ( +
+ + setLogins(event.target.value)} + placeholder="current-account, old-account" + /> +
+ )} +

+ Separate accounts with commas. Their merged changes are combined, and gateway spend is counted once +

+ {error && ( +

+ {error} +

+ )} + +
+
+
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.integration.test.tsx new file mode 100644 index 00000000000..730f0c8faf5 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.integration.test.tsx @@ -0,0 +1,198 @@ +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { afterEach, describe, expect, it, vi } from "vitest"; +import ObservedConnections from "./ObservedConnections"; +import type { ObservedSettings } from "./observedData"; + +const settings: ObservedSettings = { + id: "initial-github", + source_provider: "github", + api_url: "https://api.github.com", + repos: [], + has_token: false, + connection_type: "token", + update_interval_minutes: 1440, + ready: false, +}; +const app = { configured: true, api_url: null, callback_url: null }; +afterEach(() => vi.unstubAllGlobals()); + +describe("observed ROI connections", () => { + it("identifies the saved connection when editing its host", async () => { + const writes: unknown[] = []; + const existing = { ...settings, id: "saved-github", has_token: true, repos: ["org/service"], ready: true }; + vi.stubGlobal( + "fetch", + vi.fn(async (input: string, init: RequestInit) => { + const path = new URL(input, "http://localhost").pathname; + if (path.endsWith("/apps")) return Response.json({ github: app, gitlab: app }); + if (path.endsWith("/repositories")) return Response.json({ repositories: [], has_more: false }); + if (path.endsWith("/settings")) { + writes.push(JSON.parse(String(init.body))); + return Response.json({ ...existing, id: "enterprise-github", api_url: "https://git.example.test/api/v3" }); + } + throw new Error(path); + }), + ); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByRole("button", { name: "Edit GitHub api.github.com" })); + await user.click(screen.getByRole("button", { name: "Change connection" })); + await user.click(screen.getByText("Self-hosted instance")); + fireEvent.change(screen.getByLabelText("API URL"), { target: { value: "https://git.example.test/api/v3" } }); + fireEvent.change(screen.getByLabelText("GitHub access token"), { target: { value: "enterprise-test-token" } }); + await user.click(screen.getByRole("button", { name: "Continue" })); + expect(await screen.findByRole("heading", { name: "Choose repositories" })).toBeInTheDocument(); + expect(writes).toEqual([ + { + connection_id: "saved-github", + source_provider: "github", + api_url: "https://git.example.test/api/v3", + token: "enterprise-test-token", + repos: [], + update_interval_minutes: 1440, + }, + ]); + }); + it("starts a GitHub installation when switching from a GitLab app connection", async () => { + const starts = vi.fn(); + vi.stubGlobal( + "fetch", + vi.fn(async (input: string, init: RequestInit) => { + const url = new URL(input, "http://localhost"); + if (url.pathname.endsWith("/apps")) + return Response.json({ github: { ...app, can_install: true }, gitlab: app }); + if (url.pathname.endsWith("/repositories")) return Response.json({ repositories: [], has_more: false }); + starts(url.searchParams.get("install"), init.credentials); + return Response.json({ detail: "Authorization test stopped before redirect" }, { status: 502 }); + }), + ); + const user = userEvent.setup(); + const connected: ObservedSettings = { + ...settings, + source_provider: "gitlab", + connection_type: "app", + has_token: true, + }; + render(); + await user.click(screen.getByRole("button", { name: "Change connection" })); + await user.click(screen.getByRole("button", { name: "GitHub", exact: true })); + const connect = screen.getByRole("button", { name: "Connect GitHub", exact: true }); + await waitFor(() => expect(connect).toBeEnabled()); + await user.click(connect); + expect(await screen.findByRole("alert")).toHaveTextContent("Authorization test stopped before redirect"); + expect(starts).toHaveBeenCalledWith("true", "include"); + }); + it.each(["GitHub", "GitLab"] as const)( + "connects %s with a token, saves repositories, and starts a sync", + async (label) => { + const provider = label === "GitHub" ? "github" : "gitlab"; + const apiUrl = provider === "github" ? "https://api.github.com" : "https://gitlab.com/api/v4"; + const saved = vi.fn(); + const closed = vi.fn(); + const writes: { path: string; body: unknown }[] = []; + vi.stubGlobal( + "fetch", + vi.fn(async (input: string, init: RequestInit) => { + const path = new URL(input, "http://localhost").pathname; + if (init.method === "PUT" || init.method === "POST") + writes.push({ path, body: init.body ? JSON.parse(String(init.body)) : undefined }); + if (path.endsWith("/apps")) return Response.json({ github: app, gitlab: app }); + if (path.endsWith("/repositories")) + return Response.json({ + repositories: [{ name: "org/service", visibility: "private", archived: false }], + has_more: false, + }); + if (path.endsWith("/settings")) { + const request = JSON.parse(String(init.body)) as { repos: string[] }; + const connected = { + ...settings, + id: `saved-${provider}`, + source_provider: provider, + api_url: apiUrl, + has_token: true, + repos: request.repos, + ready: request.repos.length > 0, + }; + return Response.json(connected); + } + if (path.endsWith("/sync")) return Response.json({ running: true }, { status: 202 }); + throw new Error(path); + }), + ); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByRole("button", { name: label, exact: true })); + await user.click(screen.getByRole("button", { name: "Access token", exact: true })); + fireEvent.change(screen.getByLabelText(`${label} access token`), { target: { value: "source-test-token" } }); + await user.click(screen.getByRole("button", { name: "Continue" })); + await user.click(await screen.findByRole("checkbox", { name: /org\/service/ })); + await user.click(screen.getByRole("button", { name: "Save and sync" })); + await waitFor(() => expect(saved).toHaveBeenCalledOnce()); + expect(closed).toHaveBeenCalledOnce(); + expect(writes).toEqual([ + { + path: "/roi-calculator/observed/settings", + body: { + source_provider: provider, + api_url: apiUrl, + token: "source-test-token", + repos: [], + update_interval_minutes: 1440, + }, + }, + { + path: "/roi-calculator/observed/settings", + body: { + connection_id: `saved-${provider}`, + source_provider: provider, + api_url: apiUrl, + repos: ["org/service"], + update_interval_minutes: 1440, + }, + }, + { path: "/roi-calculator/observed/sync", body: undefined }, + ]); + }, + ); + it.each(["GitHub", "GitLab"] as const)( + "starts %s app authorization with a browser cookie and displays provider failures", + async (label) => { + const provider = label.toLowerCase(); + vi.stubGlobal( + "fetch", + vi.fn(async (input: string, init: RequestInit) => { + if (String(input).endsWith("/apps")) return Response.json({ github: app, gitlab: app }); + expect(String(input)).toContain(`/oauth/${provider}/start`); + expect(init.credentials).toBe("include"); + expect(init.method).toBe("POST"); + return Response.json({ detail: "Provider unavailable. Try again" }, { status: 502 }); + }), + ); + const user = userEvent.setup(); + render( + , + ); + await user.click(screen.getByRole("button", { name: label, exact: true })); + const button = await screen.findByRole("button", { name: `Connect ${label}`, exact: true }); + await waitFor(() => expect(button).toBeEnabled()); + await user.click(button); + expect(await screen.findByRole("alert")).toHaveTextContent("Provider unavailable. Try again"); + expect(button).toBeEnabled(); + }, + ); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.tsx new file mode 100644 index 00000000000..591b7066e48 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedConnections.tsx @@ -0,0 +1,606 @@ +"use client"; + +import { useEffect, useState } from "react"; +import { Github, Gitlab, ArrowLeft, KeyRound } from "lucide-react"; +import { z } from "zod"; +import { apiClient } from "@/components/networking"; +import { extractProxyErrorMessage } from "@/lib/http/client"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Dialog, DialogContent, DialogHeader, DialogTitle, DialogDescription } from "@/components/ui/dialog"; +import { + observedSettingsSchema, + repositoryNames, + type ObservedSettings, + type ObservedConnection, +} from "./observedData"; + +const appFields = { + configured: z.boolean(), + can_install: z.boolean().optional().default(false), + api_url: z.string().nullable(), + callback_url: z.string().nullable(), +}; +const appSchema = z.object(appFields); +const appsSchema = z.object({ github: appSchema, gitlab: appSchema }); +const repositoriesSchema = z.object({ + repositories: z.array(z.object({ name: z.string(), visibility: z.string(), archived: z.boolean() })), + has_more: z.boolean(), +}); +const defaultUrl = { github: "https://api.github.com", gitlab: "https://gitlab.com/api/v4" }; + +function preferredConnectionMethod(selected: "app" | "token" | null, configured: boolean | undefined) { + return selected ?? (configured ? "app" : "token"); +} + +function TokenFields({ + label, + provider, + token, + onToken, + apiUrl, + onUrl, + hasToken, +}: { + label: string; + provider: ObservedSettings["source_provider"]; + token: string; + onToken: (value: string) => void; + apiUrl: string; + onUrl: (value: string) => void; + hasToken: boolean; +}) { + return ( + <> + + onToken(event.target.value)} + placeholder={hasToken ? "Leave blank to keep the saved token" : "Optional for public repositories"} + /> +

+ {provider === "github" + ? "Fine-grained token: read access to pull requests, issues, and metadata" + : "Token with read_api scope"} +

+
+ Self-hosted instance + + onUrl(event.target.value)} /> +
+ + ); +} +function AppMessage({ configured, label }: { configured: boolean; label: string }) { + return ( +

+ {configured + ? `You’ll authorize ${label}, then choose repositories` + : `Register the ${label} app in gateway settings, or connect with a token`} +

+ ); +} + +function RepositoryChoices({ + available, + repos, + setRepos, + query, + setQuery, + page, + setPage, +}: { + available: z.infer | null; + repos: string; + setRepos: (value: string) => void; + query: string; + setQuery: (value: string) => void; + page: number; + setPage: (value: number) => void; +}) { + return ( + <> + { + setQuery(event.target.value); + setPage(1); + }} + placeholder="Find repositories…" + /> +
+ {!available && ( +

+ Loading repositories… +

+ )} + {available?.repositories.length === 0 && ( +

No repositories found

+ )} + {available?.repositories + .filter((repo) => !repo.archived) + .map((repo) => ( + + ))} +
+
+ + +
+ + ); +} + +function ConnectionMethod({ + method, + setMethod, + label, + provider, + token, + setToken, + apiUrl, + setApiUrl, + connected, + apps, + busy, + connect, +}: { + method: "app" | "token"; + setMethod: (value: "app" | "token") => void; + label: string; + provider: ObservedSettings["source_provider"]; + token: string; + setToken: (value: string) => void; + apiUrl: string; + setApiUrl: (value: string) => void; + connected: ObservedConnection; + apps: z.infer | null; + busy: boolean; + connect: () => void; +}) { + const connectLabel = method === "app" ? `Connect ${label}` : "Continue"; + const sameSource = connected.source_provider === provider && connected.api_url === apiUrl; + const hasSavedToken = sameSource && connected.has_token && connected.connection_type === "token"; + return ( + <> +
+ + +
+ {method === "token" && ( + + )} + {method === "app" && } + + + ); +} + +type ConnectionStep = "list" | "connect" | "repos"; +const stepTitles = { list: "Connections", connect: "Connect your code", repos: "Choose repositories" }; + +function initialStep(settings: ObservedSettings, afterAuthorization: boolean): ConnectionStep { + if (afterAuthorization) return "repos"; + if (settings.connections?.length) return "list"; + return settings.has_token || settings.ready ? "repos" : "connect"; +} + +function stepDescription(step: ConnectionStep, label: string) { + if (step === "list") return "All selected repositories appear in one report"; + if (step === "connect") return "Connect GitHub and GitLab with an app or access token"; + return `Select ${label} repositories to compare`; +} + +function connectionMethodLabel(entry: ObservedConnection) { + if (entry.connection_type === "app") return "App"; + return entry.has_token ? "Token" : "Public access"; +} + +function ConnectionList({ + connections, + onEdit, + onAdd, +}: { + connections: ObservedConnection[]; + onEdit: (entry: ObservedConnection) => void; + onAdd: () => void; +}) { + return ( + <> + {connections.map((entry) => ( +
+
+

+ {entry.source_provider === "github" ? : } + {entry.source_provider === "github" ? "GitHub" : "GitLab"} +

+

{new URL(entry.api_url).host}

+

+ {entry.repos.length} repositories · {connectionMethodLabel(entry)} +

+
+ +
+ ))} + + + ); +} + +function initialMethod(settings: ObservedSettings) { + return settings.has_token ? settings.connection_type : null; +} + +function hasConnections(settings: ObservedSettings) { + return Boolean(settings.connections?.length); +} + +function canManageApp(connected: ObservedConnection, apps: z.infer | null) { + return ( + connected.connection_type === "app" && connected.source_provider === "github" && Boolean(apps?.github.can_install) + ); +} + +export default function ObservedConnections({ + accessToken, + settings, + onClose, + onSaved, + initialError = "", + afterAuthorization = false, +}: { + accessToken: string; + settings: ObservedSettings; + onClose: () => void; + onSaved: () => void; + initialError?: string; + afterAuthorization?: boolean; +}) { + const [savedSettings, setSavedSettings] = useState(settings); + const [connected, setConnected] = useState(settings); + const [provider, setProvider] = useState(settings.source_provider); + const [apiUrl, setApiUrl] = useState(settings.api_url); + const [selectedMethod, setMethod] = useState<"app" | "token" | null>(initialMethod(settings)); + const [step, setStep] = useState(() => initialStep(settings, afterAuthorization)); + const [token, setToken] = useState(""); + const [repos, setRepos] = useState(settings.repos.join(", ")); + const [apps, setApps] = useState | null>(null); + const method = preferredConnectionMethod(selectedMethod, apps?.[provider].configured); + const [available, setAvailable] = useState | null>(null); + const [query, setQuery] = useState(""); + const [page, setPage] = useState(1); + const [busy, setBusy] = useState(false); + const [error, setError] = useState(initialError); + const label = provider === "github" ? "GitHub" : "GitLab"; + const saveLabel = repositoryNames(repos).length ? "Save and sync" : "Save repositories"; + const manageApp = canManageApp(connected, apps); + const showConnections = hasConnections(savedSettings); + useEffect(() => { + const controller = new AbortController(); + apiClient + .get("/roi-calculator/observed/apps", { accessToken, signal: controller.signal }) + .then((data) => { + if (!controller.signal.aborted) setApps(appsSchema.parse(data)); + }) + .catch((reason: unknown) => { + if (!controller.signal.aborted) setError(extractProxyErrorMessage(reason)); + }); + return () => controller.abort(); + }, [accessToken]); + useEffect(() => { + if (step !== "repos" || (!connected.has_token && connected.source_provider === "github")) return; + const controller = new AbortController(); + const timer = setTimeout(() => { + apiClient + .get("/roi-calculator/observed/repositories", { + accessToken, + signal: controller.signal, + query: { query, page, connection: connected.id }, + }) + .then((data) => { + if (!controller.signal.aborted) setAvailable(repositoriesSchema.parse(data)); + }) + .catch((reason: unknown) => { + if (!controller.signal.aborted) setError(extractProxyErrorMessage(reason)); + }); + }, 250); + return () => { + clearTimeout(timer); + controller.abort(); + }; + }, [accessToken, step, connected, query, page]); + function selectProvider(value: ObservedSettings["source_provider"]) { + setProvider(value); + const existing = savedSettings.connections?.find( + (entry) => entry.source_provider === value && entry.api_url === defaultUrl[value], + ); + setApiUrl(existing?.api_url ?? defaultUrl[value]); + setConnected( + existing ?? { + ...settings, + source_provider: value, + api_url: defaultUrl[value], + repos: [], + has_token: false, + ready: false, + connection_type: "token", + id: undefined, + }, + ); + setMethod(existing?.connection_type ?? null); + setToken(""); + setError(""); + } + function searchRepositories(value: string) { + setAvailable(null); + setError(""); + setQuery(value); + } + function changePage(value: number) { + setAvailable(null); + setError(""); + setPage(value); + } + async function connect(install = false) { + setBusy(true); + setError(""); + try { + if (method === "app" || install) { + const sameApp = connected.source_provider === provider && connected.connection_type === "app"; + const firstInstallation = provider === "github" && !sameApp && apps?.github.can_install; + const result = z.object({ url: z.string().url() }).parse( + await apiClient.post(`/roi-calculator/observed/oauth/${provider}/start`, { + accessToken, + credentials: "include", + query: { install: install || Boolean(firstInstallation) }, + }), + ); + window.location.assign(result.url); + return; + } + const same = provider === connected.source_provider && apiUrl === connected.api_url; + const keepToken = same && connected.has_token && connected.connection_type === "token"; + const result = observedSettingsSchema.parse( + await apiClient.put("/roi-calculator/observed/settings", { + accessToken, + body: { + connection_id: savedSettings.connections?.find((entry) => entry.id === connected.id)?.id, + source_provider: provider, + api_url: apiUrl, + token: token || (keepToken ? undefined : ""), + repos: same ? connected.repos : [], + update_interval_minutes: connected.update_interval_minutes, + }, + }), + ); + setSavedSettings(result); + setConnected(result); + setToken(""); + setRepos(result.repos.join(", ")); + setAvailable(null); + setStep("repos"); + } catch (reason) { + setError(extractProxyErrorMessage(reason)); + } finally { + setBusy(false); + } + } + async function save() { + setBusy(true); + setError(""); + try { + const result = observedSettingsSchema.parse( + await apiClient.put("/roi-calculator/observed/settings", { + accessToken, + body: { + connection_id: connected.id, + source_provider: connected.source_provider, + api_url: connected.api_url, + repos: repositoryNames(repos), + update_interval_minutes: connected.update_interval_minutes, + }, + }), + ); + if (result.ready) await apiClient.post("/roi-calculator/observed/sync", { accessToken }); + onSaved(); + onClose(); + } catch (reason) { + setError(extractProxyErrorMessage(reason)); + } finally { + setBusy(false); + } + } + function edit(entry: ObservedConnection) { + setConnected(entry); + setProvider(entry.source_provider); + setApiUrl(entry.api_url); + setMethod(entry.connection_type); + setRepos(entry.repos.join(", ")); + setToken(""); + setQuery(""); + setPage(1); + setAvailable(null); + setError(""); + setStep("repos"); + } + return ( + { + if (!open) onClose(); + }} + > + + + {stepTitles[step]} + {stepDescription(step, label)} + +
+ {step === "list" && ( + { + selectProvider( + savedSettings.connections?.some((entry) => entry.source_provider === "github") ? "gitlab" : "github", + ); + setStep("connect"); + }} + /> + )} + {step === "connect" && ( + <> + {showConnections && ( + + )} +
+ {(["github", "gitlab"] as const).map((value) => ( + + ))} +
+ connect()} + /> + + )} + {step === "repos" && ( + <> + {showConnections && ( + + )} + + {manageApp && ( + + )} + + setRepos(event.target.value)} + placeholder={ + provider === "github" ? "owner/repo, owner/another-repo" : "group/project, group/subgroup/project" + } + /> + {(connected.has_token || connected.source_provider === "gitlab") && ( + <> + + + )} + + + )} + {error && ( +

+ {error} +

+ )} +
+
+
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx new file mode 100644 index 00000000000..02d60a05029 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx @@ -0,0 +1,232 @@ +"use client"; + +import { useState } from "react"; +import { ExternalLink, GitPullRequest } from "lucide-react"; +import { Badge } from "@/components/ui/badge"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Sheet, SheetContent, SheetHeader, SheetTitle, SheetDescription } from "@/components/ui/sheet"; +import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; +import { + dateRange, + changeTerms, + duration, + money, + number, + recordedBranches, + type Comparison, + type ObservedPerson, + type ObservedPull, + type ObservedSnapshot, + type Period, +} from "./observedData"; + +export function PullList({ + pulls, + provider, +}: { + pulls: ObservedPull[]; + provider: ObservedSnapshot["source_provider"]; +}) { + const terms = changeTerms(provider); + const [query, setQuery] = useState(""); + const [limit, setLimit] = useState(20); + const filtered = pulls.filter((pull) => + `${pull.number} ${pull.title} ${pull.author} ${pull.repo} ${pull.source_repo} ${pull.source_branch}` + .toLowerCase() + .includes(query.toLowerCase()), + ); + return ( +
+ { + setQuery(event.target.value); + setLimit(20); + }} + className="max-w-sm" + /> + + + + {terms.requests} + Author + Opened to merged + Tagged spend + + + + {filtered.slice(0, limit).map((pull) => ( + + + + + + #{pull.number} + {pull.title} + {pull.repo} + + + + + + {pull.author || "Deleted author"} + {pull.agent && ( + + Agent + + )} + + {duration(pull.merge_hours)} + + {money(pull.branch_cost.spend)} + + + ))} + +
+

+ Elapsed time from opening to merge, not engineering effort or time saved +

+ {filtered.length === 0 && ( +

+ {query ? `No ${terms.lower} match this search` : `No ${terms.lower} in this period`} +

+ )} +
+ + {number(Math.min(limit, filtered.length))} of {number(filtered.length)} {terms.lower} + + {limit < filtered.length && ( + + )} +
+
+ ); +} + +export function PersonDetails({ + person, + snapshot, + comparison, + onClose, + onEdit, +}: { + person: ObservedPerson; + snapshot: ObservedSnapshot; + comparison: Comparison; + onClose: () => void; + onEdit?: () => void; +}) { + const terms = changeTerms(snapshot.source_provider); + const [period, setPeriod] = useState("current"); + const current = person.periods.current; + const baseline = person.periods[comparison]; + const urls = new Set(person.periods[period].pr_urls); + const pulls = snapshot.pulls[period].filter((pull) => urls.has(pull.url)); + return ( + { + if (!open) onClose(); + }} + > + + + {person.name} + + {person.email} · {person.logins.join(", ")} + + + {onEdit && ( + + )} +
+
+

Merged {terms.plural}

+

{number(current.merged_prs)}

+

{number(baseline.merged_prs)} in comparison

+
+
+

Recorded spend

+

+ {money(current.spend_observation === "no_records" ? null : current.gateway_recorded_spend)} +

+

Gateway only

+
+
+

Spend / matched {terms.singular}

+

{money(current.recorded_spend_per_attributed_pr)}

+

Period average

+
+
+
+

Merged {terms.plural}

+
+ + +
+
+

{dateRange(snapshot.periods[period].window)} · UTC

+ +
+
+ ); +} + +export function BranchSpend({ snapshot }: { snapshot: ObservedSnapshot }) { + const rows = recordedBranches(snapshot); + return ( +
+ + + + Repository + Branch + Requests + Tagged spend + + + + {rows.map((row) => ( + + {row.repo} + {row.branch} + {number(row.requests)} + {money(row.spend)} + + ))} + +
+ {rows.length === 0 && ( +

No tagged branch spend in this period

+ )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx new file mode 100644 index 00000000000..baab6eee876 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.integration.test.tsx @@ -0,0 +1,305 @@ +import { fireEvent, render, screen, waitFor, within } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { afterEach, describe, expect, it, vi } from "vitest"; +import ObservedROIView from "./ObservedROIView"; +import type { ObservedSettings, ObservedSnapshot, ObservedStatus } from "./observedData"; + +const settings: ObservedSettings = { + source_provider: "gitlab", + api_url: "https://gitlab.com/api/v4", + repos: ["org/service"], + has_token: true, + connection_type: "app", + update_interval_minutes: 1440, + ready: true, +}; +const idle: ObservedStatus = { + running: false, + phase: "complete", + stage: "", + done: 0, + total: 0, + error: null, + finished_at: null, +}; +const period = { + window: { start: "2026-09-01", end: "2026-09-28" }, + merged_prs: 1, + median_merge_hours: 16 / 3600, + human_authored: 1, + agent_authored: 0, + missing_author: 0, + agents_without_requester: 0, + matched_internal_prs: 0, + new_bug_labeled_issues: 0, + new_regression_labeled_issues: 0, + explicitly_titled_revert_prs: 0, + matched_users_recorded_spend: 0, + spend_observation: "no_records" as const, + human_summary: { median_merge_hours: 16 / 3600 }, +}; +const report: ObservedSnapshot = { + source_provider: "gitlab", + repos: ["org/service"], + unmatched_logins: [], + unlinked_branches: [], + captured_at: "2026-09-29T00:00:00Z", + periods: { current: period, previous: period, last_year: period }, + people: [], + pulls: { current: [], previous: [], last_year: [] }, +}; +const app = { configured: true, api_url: null, callback_url: null }; + +afterEach(() => { + vi.unstubAllGlobals(); + window.history.replaceState(null, "", "/"); +}); + +describe("observed ROI dashboard", () => { + it("previews every sample view before setup, changes sample periods without writes, and exits back to setup", async () => { + const requests = vi.fn(async (input: string, _init: RequestInit) => { + const path = new URL(input, "http://localhost").pathname; + const disconnected = { ...settings, ready: false, repos: [], has_token: false }; + if (path.endsWith("/settings")) return Response.json(disconnected); + if (path.endsWith("/report")) return Response.json({ report: null }); + if (path.endsWith("/sync")) return Response.json(idle); + throw new Error(path); + }); + vi.stubGlobal("fetch", requests); + const user = userEvent.setup(); + render(); + expect(await screen.findByRole("heading", { name: "Connect your repositories" })).toBeInTheDocument(); + expect(screen.queryByText("Ready to sync")).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "Preview sample report" })); + expect(screen.getByRole("status")).toHaveTextContent("You’re viewing demo data"); + expect(screen.getByRole("tab", { name: "Engineers 3", selected: true })).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Connections" })).not.toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Link accounts" })).not.toBeInTheDocument(); + expect(window.location.search).toBe("?demo=1"); + await user.click(screen.getByRole("button", { name: "View Alex Rivera's merged changes" })); + expect(await screen.findByRole("dialog", { name: "Alex Rivera" })).toHaveTextContent("alex-demo@example.com"); + expect(screen.getByRole("heading", { name: "Merged changes" })).toBeInTheDocument(); + expect(screen.queryByRole("button", { name: "Edit linked accounts" })).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "Close" })); + await user.click(screen.getByRole("tab", { name: "Merged changes" })); + expect(screen.getByRole("img", { name: /Merged changes by week/ })).toBeInTheDocument(); + await user.click(screen.getByRole("tab", { name: "Quality" })); + expect(screen.getByText("New regression-labeled issues")).toBeInTheDocument(); + await user.click(screen.getByRole("tab", { name: "Branch spend" })); + expect(screen.getAllByText(/feature\/sample-/).length).toBeGreaterThan(0); + await user.click(screen.getByRole("combobox", { name: "Reporting period" })); + await user.click(await screen.findByRole("option", { name: "Last 7 days" })); + expect(screen.getByRole("combobox", { name: "Reporting period" })).toHaveTextContent("Last 7 days"); + await user.click(screen.getByRole("combobox", { name: "Comparison period" })); + await user.click(await screen.findByRole("option", { name: "vs. same period last year" })); + expect(screen.getByRole("combobox", { name: "Comparison period" })).toHaveTextContent("vs. same period last year"); + await user.click(screen.getByRole("button", { name: "Exit demo" })); + expect(screen.getByRole("heading", { name: "Connect your repositories" })).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Connect GitHub or GitLab" })).toBeEnabled(); + expect(window.location.search).toBe(""); + expect(requests.mock.calls.every(([, init]) => init.method === "GET")).toBe(true); + }); + + it("keeps live sync and its report intact when entering and exiting the demo", async () => { + const requests = vi.fn(async (input: string, _init: RequestInit) => { + const path = new URL(input, "http://localhost").pathname; + if (path.endsWith("/settings")) return Response.json(settings); + if (path.endsWith("/report")) return Response.json({ report }); + if (path.endsWith("/sync")) return Response.json({ ...idle, running: true, stage: "Reading changes" }); + throw new Error(path); + }); + vi.stubGlobal("fetch", requests); + window.history.replaceState(null, "", "/roi-calculator/?from=review#report"); + const user = userEvent.setup(); + render(); + expect(await screen.findByRole("button", { name: "Cancel sync" })).toBeEnabled(); + await user.click(screen.getByRole("button", { name: "Preview sample report" })); + expect(screen.queryByRole("button", { name: "Cancel sync" })).not.toBeInTheDocument(); + expect(screen.getByRole("button", { name: "2 repositories" })).toBeInTheDocument(); + expect(window.location.search).toBe("?from=review&demo=1"); + await user.click(screen.getByRole("button", { name: "Exit demo" })); + expect(screen.getByRole("button", { name: "Cancel sync" })).toBeEnabled(); + expect(screen.getByRole("button", { name: "1 repository" })).toBeInTheDocument(); + expect(screen.getByRole("tab", { name: "Merge requests", selected: true })).toBeInTheDocument(); + expect(window.location.search).toBe("?from=review"); + expect(window.location.hash).toBe("#report"); + expect(requests.mock.calls.every(([, init]) => init.method === "GET")).toBe(true); + }); + + it.each(["failed", "pending"])("opens a demo URL even when live requests are %s", async (state) => { + window.history.replaceState(null, "", "/roi-calculator/?demo=1"); + const pending = Promise.withResolvers(); + vi.stubGlobal( + "fetch", + vi.fn(() => (state === "failed" ? Promise.reject(new Error("Live data unavailable")) : pending.promise)), + ); + const user = userEvent.setup(); + render(); + expect(screen.getByRole("status")).toHaveTextContent("You’re viewing demo data"); + expect(screen.getByText("Alex Rivera")).toBeInTheDocument(); + expect(screen.queryByRole("alert")).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "Exit demo" })); + expect(screen.queryByText("Alex Rivera")).not.toBeInTheDocument(); + expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); + expect(window.location.search).toBe(""); + if (state === "failed") expect(await screen.findByRole("alert")).toHaveTextContent("Live data unavailable"); + }); + + it("retries failures, keeps the report during cancellation, and refreshes after completion", async () => { + let status: ObservedStatus = { ...idle, phase: "error", error: "Provider temporarily unavailable" }; + let completeOnPoll = false; + let currentReport = { ...report, repos: ["org/service", "org/docs"] }; + vi.stubGlobal( + "fetch", + vi.fn(async (input: string, init: RequestInit) => { + const path = new URL(input, "http://localhost").pathname; + if (path.endsWith("/settings")) return Response.json(settings); + if (path.endsWith("/report")) return Response.json({ report: currentReport }); + if (path.endsWith("/sync")) { + if (init.method === "POST") status = { ...idle, running: true, phase: "pulls" }; + if (init.method === "DELETE") status = { ...idle, phase: "cancelled" }; + if (init.method === "GET" && completeOnPoll) { + status = { ...idle, finished_at: "2026-09-29T00:01:00Z" }; + currentReport = { ...report, repos: ["org/updated"] }; + } + return Response.json(status); + } + throw new Error(path); + }), + ); + const user = userEvent.setup(); + render(); + expect(await screen.findByRole("tab", { name: "Merge requests", selected: true })).toBeInTheDocument(); + await user.click(screen.getByRole("tab", { name: "Quality" })); + expect(screen.getByRole("alert")).toHaveTextContent("Provider temporarily unavailable"); + await user.click(screen.getByRole("button", { name: "Retry" })); + await user.click(await screen.findByRole("button", { name: "Cancel sync" })); + expect(await screen.findByRole("button", { name: "Sync now" })).toBeEnabled(); + expect(screen.queryByText("org/service")).not.toBeInTheDocument(); + await user.click(screen.getByRole("button", { name: "2 repositories" })); + const repositories = await screen.findByRole("dialog", { name: "Repositories" }); + expect(within(repositories).getByText("org/service")).toBeInTheDocument(); + expect(within(repositories).getByText("org/docs")).toBeInTheDocument(); + await user.keyboard("{Escape}"); + await waitFor(() => expect(screen.queryByRole("dialog", { name: "Repositories" })).not.toBeInTheDocument()); + await user.click(screen.getByRole("button", { name: "Sync now" })); + expect(await screen.findByRole("button", { name: "Cancel sync" })).toBeEnabled(); + completeOnPoll = true; + await user.click(await screen.findByRole("button", { name: "1 repository" }, { timeout: 4000 })); + expect( + await within(screen.getByRole("dialog", { name: "Repositories" })).findByText("org/updated"), + ).toBeInTheDocument(); + await user.keyboard("{Escape}"); + expect(screen.getByRole("tab", { name: "Quality", selected: true })).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Sync now" })).toBeEnabled(); + }); + + it("syncs the selected range and labels its equal-length comparison", async () => { + let currentReport = report; + const requested: string[] = []; + vi.stubGlobal( + "fetch", + vi.fn(async (input: string, init: RequestInit) => { + const url = new URL(input, "http://localhost"); + if (url.pathname.endsWith("/settings")) return Response.json(settings); + if (url.pathname.endsWith("/report")) return Response.json({ report: currentReport }); + if (url.pathname.endsWith("/sync")) { + if (init.method === "POST") { + requested.push(url.searchParams.get("days") ?? ""); + currentReport = { + ...report, + periods: { + ...report.periods, + current: { ...period, window: { start: "2026-09-22", end: "2026-09-28" } }, + previous: { ...period, window: { start: "2026-09-15", end: "2026-09-21" } }, + }, + }; + } + return Response.json(idle); + } + throw new Error(url.pathname); + }), + ); + const user = userEvent.setup(); + render(); + await user.click(await screen.findByRole("combobox", { name: "Reporting period" })); + await user.click(await screen.findByRole("option", { name: "Last 7 days" })); + await waitFor(() => expect(requested).toEqual(["7"])); + expect(await screen.findByText(/Comparing with Sep 15.*Sep 21/)).toBeInTheDocument(); + expect(screen.getByRole("combobox", { name: "Reporting period" })).toHaveTextContent("Last 7 days"); + expect(screen.getByRole("combobox", { name: "Comparison period" })).toHaveTextContent("vs. previous period"); + }); + + it("shows a successful empty repository without a setup prompt or invented durations", async () => { + const emptyPeriod = { + ...period, + merged_prs: 0, + human_authored: 0, + median_merge_hours: null, + human_summary: { median_merge_hours: null }, + }; + const empty = { ...report, periods: { current: emptyPeriod, previous: emptyPeriod, last_year: emptyPeriod } }; + vi.stubGlobal( + "fetch", + vi.fn(async (input: string) => { + const path = new URL(input, "http://localhost").pathname; + if (path.endsWith("/settings")) return Response.json(settings); + if (path.endsWith("/report")) return Response.json({ report: empty }); + if (path.endsWith("/sync")) return Response.json(idle); + throw new Error(path); + }), + ); + render(); + expect(await screen.findByRole("heading", { name: "No merged changes yet" })).toBeInTheDocument(); + expect(screen.getByText("No merges")).toBeInTheDocument(); + expect(screen.getByRole("button", { name: "Sync now" })).toBeEnabled(); + expect(screen.queryByRole("alert")).not.toBeInTheDocument(); + expect(screen.queryByRole("heading", { name: "Connect your repositories" })).not.toBeInTheDocument(); + expect(screen.queryByText("0h")).not.toBeInTheDocument(); + }); + + it.each([ + { query: "connected=gitlab", alerts: [] }, + { query: "connection_failed=1", alerts: ["Connection failed or expired. Try again or use a token"] }, + { query: "connection_cancelled=1", alerts: ["Connection cancelled. Choose an app or token to try again"] }, + ])("resumes setup after $query and refreshes saved changes after closing", async ({ query, alerts }) => { + window.history.replaceState(null, "", `/roi-calculator/?${query}`); + let currentSettings = settings; + vi.stubGlobal( + "fetch", + vi.fn(async (input: string, init: RequestInit) => { + const path = new URL(input, "http://localhost").pathname; + if (path.endsWith("/apps")) return Response.json({ github: app, gitlab: app }); + if (path.endsWith("/repositories")) return Response.json({ repositories: [], has_more: false }); + if (path.endsWith("/settings")) { + if (init.method === "PUT") currentSettings = { ...settings, repos: ["org/changed"] }; + return Response.json(currentSettings); + } + if (path.endsWith("/report")) return Response.json({ report: { ...report, repos: currentSettings.repos } }); + if (path.endsWith("/sync")) + return init.method === "POST" + ? Response.json({ detail: "Provider unavailable" }, { status: 502 }) + : Response.json(idle); + throw new Error(path); + }), + ); + const user = userEvent.setup(); + render(); + const dialog = await screen.findByRole("dialog", { name: "Choose repositories" }); + expect( + within(dialog) + .queryAllByRole("alert") + .map((alert) => alert.textContent), + ).toEqual(alerts); + expect(window.location.search).toBe(""); + fireEvent.change(within(dialog).getByLabelText("Repositories"), { target: { value: "org/changed" } }); + await user.click(within(dialog).getByRole("button", { name: "Save and sync" })); + expect(await within(dialog).findByRole("alert")).toHaveTextContent("Provider unavailable"); + await user.click(within(dialog).getByRole("button", { name: "Close" })); + await waitFor(() => expect(screen.queryByRole("dialog")).not.toBeInTheDocument()); + await user.click(await screen.findByRole("button", { name: "1 repository" })); + expect( + await within(screen.getByRole("dialog", { name: "Repositories" })).findByText("org/changed"), + ).toBeInTheDocument(); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx new file mode 100644 index 00000000000..29199b67f22 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedROIView.tsx @@ -0,0 +1,253 @@ +"use client"; + +import { useEffect, useState } from "react"; +import { Link2, RefreshCw } from "lucide-react"; +import { apiClient } from "@/components/networking"; +import { extractProxyErrorMessage } from "@/lib/http/client"; +import { Page } from "@/components/shared/Page"; +import { PageHeader, PageHeaderDescription, PageHeaderTitle } from "@/components/shared/PageHeader"; +import { DemoNotice } from "@/components/shared/DemoNotice"; +import { Button } from "@/components/ui/button"; +import { Skeleton } from "@/components/ui/skeleton"; +import ObservedConnections from "./ObservedConnections"; +import ObservedReport from "./ObservedReport"; +import { useObservedReport, type ObservedViewData } from "./useObservedReport"; +import { syncMessage, type ObservedSnapshot } from "./observedData"; +import { createObservedDemo } from "./observedDemo"; + +function SyncActions({ + data, + error, + busy, + readOnly, + compact = false, + onSync, + onRetry, +}: { + data: ObservedViewData | null; + error: string; + busy: boolean; + readOnly: boolean; + compact?: boolean; + onSync: (cancel: boolean) => void; + onRetry: () => void; +}) { + const message = error || data?.status.error; + const statusMessage = data ? syncMessage(data.status, data.report) : ""; + const canSync = data?.settings.ready && !readOnly; + if (!message && !statusMessage && !canSync) return null; + return ( +
+ {message && ( +
+ {message} + +
+ )} + {data && ( +
+ {statusMessage && ( + + {statusMessage} + + )} + {!readOnly && data.settings.ready && ( + + )} +
+ )} +
+ ); +} + +function EmptyReport({ + data, + readOnly, + onConnect, +}: { + data: ObservedViewData; + readOnly: boolean; + onConnect: () => void; +}) { + function title() { + if (data.status.running) return "Reading repository activity"; + return data.settings.ready ? "Ready for your first report" : "Connect your repositories"; + } + return ( +
+

{title()}

+

+ {data.status.running + ? "Your report will appear here when the first sync finishes" + : "Compare merged changes, issue trends, and recorded AI spend across your team"} +

+ {!readOnly && !data.status.running && ( + + )} +
+ ); +} + +export default function ObservedROIView({ + accessToken, + isViewOnly = false, +}: { + accessToken: string; + isViewOnly?: boolean; +}) { + const { data, error, refresh } = useObservedReport(accessToken); + const [returned] = useState(() => new URLSearchParams(typeof window === "undefined" ? "" : window.location.search)); + const [sample, setSample] = useState(() => + returned.get("demo") === "1" ? createObservedDemo(28) : null, + ); + const [connections, setConnections] = useState( + ["github", "gitlab"].includes(returned.get("connected") ?? "") || + returned.has("connection_cancelled") || + returned.has("connection_failed"), + ); + const [connectionError, setConnectionError] = useState(() => { + if (returned.has("connection_failed")) return "Connection failed or expired. Try again or use a token"; + if (returned.has("connection_cancelled")) return "Connection cancelled. Choose an app or token to try again"; + return ""; + }); + const [afterAuthorization, setAfterAuthorization] = useState(Boolean(returned.get("connected"))); + const [busy, setBusy] = useState(false); + const [actionError, setActionError] = useState(""); + useEffect(() => { + const url = new URL(window.location.href); + url.searchParams.delete("connected"); + url.searchParams.delete("connection_cancelled"); + url.searchParams.delete("connection_failed"); + window.history.replaceState(window.history.state, "", url); + }, []); + function previewSample(enabled: boolean) { + const url = new URL(window.location.href); + if (enabled) url.searchParams.set("demo", "1"); + else url.searchParams.delete("demo"); + window.history.replaceState(window.history.state, "", url); + setSample(enabled ? createObservedDemo(28) : null); + } + function closeConnections() { + setConnections(false); + setConnectionError(""); + setAfterAuthorization(false); + refresh(); + } + async function sync(cancel: boolean, days?: number) { + setBusy(true); + setActionError(""); + try { + if (cancel) await apiClient.delete("/roi-calculator/observed/sync", { accessToken }); + else await apiClient.post("/roi-calculator/observed/sync", { accessToken, query: { days } }); + refresh(); + } catch (reason) { + setActionError(extractProxyErrorMessage(reason)); + } finally { + setBusy(false); + } + } + function retry() { + if (error || isViewOnly || !data?.settings.ready) refresh(); + else void sync(data.status.running); + } + if (sample) { + return ( + setConnections(true)} + actions={null} + notice={ previewSample(false)} />} + syncing={false} + onPeriod={(days) => setSample(createObservedDemo(days))} + /> + ); + } + const previewButton = ( + + ); + const actions = ( + <> + {previewButton} + + + ); + const content = data?.report ? ( + setConnections(true)} + actions={actions} + syncing={busy || data.status.running} + onPeriod={isViewOnly ? undefined : (days) => void sync(false, days)} + /> + ) : ( + + +
+ ROI Calculator + {previewButton} +
+ Are we shipping more, with fewer bugs, at a better cost? +
+ + {!data && !error && ( + <> + + + + )} + {data && setConnections(true)} />} +
+ ); + const showConnections = connections && data && !isViewOnly; + return ( + <> + {content} + {showConnections && ( + + )} + + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedReport.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedReport.tsx new file mode 100644 index 00000000000..d68750480cc --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedReport.tsx @@ -0,0 +1,611 @@ +"use client"; + +import { useState } from "react"; +import { ArrowDown, ArrowUp, CalendarDays, ChevronDown, ChevronRight, Link2, Search, Users } from "lucide-react"; +import { Page, PageTabsList, PageTabsTrigger } from "@/components/shared/Page"; +import { PageHeader, PageHeaderTitle } from "@/components/shared/PageHeader"; +import { Button } from "@/components/ui/button"; +import { Input } from "@/components/ui/input"; +import { Popover, PopoverContent, PopoverTitle, PopoverTrigger } from "@/components/ui/popover"; +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { Tabs, TabsContent } from "@/components/ui/tabs"; +import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; +import ObservedAccounts from "./ObservedAccounts"; +import { BranchSpend, PersonDetails, PullList } from "./ObservedDetails"; +import { + change, + dateRange, + changeTerms, + duration, + money, + number, + visiblePeople, + weeklyMerges, + type Comparison, + type ObservedPerson, + type ObservedSnapshot, + type PeopleSort, +} from "./observedData"; + +function Delta({ + current, + baseline, + neutral = false, +}: { + current: number | null; + baseline: number | null; + neutral?: boolean; +}) { + const delta = current === null || baseline === null ? null : change(current, baseline); + if (delta === null) return No baseline; + const Icon = delta >= 0 ? ArrowUp : ArrowDown; + return ( + = 0 ? "increase" : "decrease"}`} + className={`inline-flex items-center gap-1 text-xs tabular-nums ${neutral ? "text-muted-foreground" : "text-foreground"}`} + > + + {number(Math.abs(delta))}% + + ); +} + +function Metric({ + label, + value, + detail, + current, + baseline, +}: { + label: string; + value: string; + detail: string; + current?: number | null; + baseline?: number | null; +}) { + return ( +
+
{label}
+
+ {value} + {current !== undefined && baseline !== undefined && } +
+

{detail}

+
+ ); +} + +function ShippingTrend({ snapshot, comparison }: { snapshot: ObservedSnapshot; comparison: Comparison }) { + const terms = changeTerms(snapshot.source_provider); + const current = weeklyMerges(snapshot, "current"); + const baseline = weeklyMerges(snapshot, comparison); + const max = Math.max(1, ...current, ...baseline); + return ( +
+
+

Shipping activity

+
+ + + Current period + + + + {comparison === "previous" ? "Previous period" : "Last year"} + +
+
+
+ {current.map((value, week) => ( +
+
+
+ + {baseline[week]} + +
+
+ {value} +
+
+

W{week + 1}

+
+ ))} +
+
+ ); +} + +function PeopleTable({ + snapshot, + comparison, + onSelect, +}: { + snapshot: ObservedSnapshot; + comparison: Comparison; + onSelect: (person: ObservedPerson) => void; +}) { + const terms = changeTerms(snapshot.source_provider); + const [query, setQuery] = useState(""); + const [sort, setSort] = useState("merged"); + const people = visiblePeople(snapshot.people, query, sort); + return ( +
+
+
+ + setQuery(event.target.value)} + className="pl-9" + /> +
+
+ + {people.length} {people.length === 1 ? "engineer" : "engineers"} + + +
+
+
+ + + + Engineer + Merged {terms.plural} + Authored / agent + + {comparison === "previous" ? "vs. previous" : "vs. last year"} + + Median merge + Recorded spend + Spend / {terms.singular} + + Details + + + + + {people.map((person) => { + const current = person.periods.current; + const baseline = person.periods[comparison]; + return ( + + + + + {number(current.merged_prs)} + +
+
+ + +
+ + {current.direct_authored}/{current.declared_agent_owned} + +
+
+ + + ({baseline.merged_prs}) + + {duration(current.median_merge_hours)} + + {money(current.spend_observation === "no_records" ? null : current.gateway_recorded_spend)} + + + {money(current.recorded_spend_per_attributed_pr)} + + + + +
+ ); + })} +
+
+ {people.length === 0 && ( +
+ {query ? `No engineers match “${query}”` : "Link accounts to see your engineers"} +
+ )} +
+

+ + + Authored + + + + Agent, explicit requester + + Spend / {terms.singular} is recorded period spend divided by matched {terms.plural} +

+
+ ); +} + +function Quality({ snapshot, comparison }: { snapshot: ObservedSnapshot; comparison: Comparison }) { + const terms = changeTerms(snapshot.source_provider); + const current = snapshot.periods.current; + const baseline = snapshot.periods[comparison]; + const rows = [ + { + label: "New bug-labeled issues", + current: current.new_bug_labeled_issues, + baseline: baseline.new_bug_labeled_issues, + detail: "Opened during the period, with bug or kind:bug labels at collection", + }, + { + label: "New regression-labeled issues", + current: current.new_regression_labeled_issues, + baseline: baseline.new_regression_labeled_issues, + detail: "Opened during the period and labeled as regressions", + }, + { + label: `Revert-titled ${terms.plural}`, + current: current.explicitly_titled_revert_prs, + baseline: baseline.explicitly_titled_revert_prs, + detail: `Merged ${terms.plural} whose titles explicitly indicate a revert`, + }, + ]; + return ( +
+
+ + + + Repository signal + Current + Comparison + Change + + + + {rows.map((row) => ( + + +

{row.label}

+

{row.detail}

+
+ {number(row.current)} + {number(row.baseline)} + + + +
+ ))} +
+
+
+

+ These signals help check whether more shipping comes with more bugs. Labels and revert titles are incomplete + proxies; they do not establish a change-failure rate or attribute bugs to an engineer. +

+
+ ); +} + +function costPerChange(period: ObservedSnapshot["periods"]["current"]) { + if (period.spend_observation !== "records_present" || period.matched_internal_prs === 0) return null; + return period.matched_users_recorded_spend / period.matched_internal_prs; +} + +export default function ObservedReport({ + snapshot, + accessToken, + readOnly, + onRefresh, + onConnect, + actions, + notice, + syncing, + onPeriod, +}: { + snapshot: ObservedSnapshot; + accessToken: string; + readOnly: boolean; + onRefresh: () => void; + onConnect: () => void; + actions: React.ReactNode; + notice?: React.ReactNode; + syncing: boolean; + onPeriod?: (days: number) => void; +}) { + const [comparison, setComparison] = useState("previous"); + const [activeTab, setActiveTab] = useState( + snapshot.people.length && snapshot.periods.current.merged_prs > 0 ? "people" : "pulls", + ); + const [accountEmail, setAccountEmail] = useState(null); + const [personEmail, setPersonEmail] = useState(null); + const person = snapshot.people.find((entry) => entry.email === personEmail) ?? null; + const terms = changeTerms(snapshot.source_provider); + const current = snapshot.periods.current; + const baseline = snapshot.periods[comparison]; + const days = Math.round((Date.parse(current.window.end) - Date.parse(current.window.start)) / 86400000) + 1; + const rangeOptions = [...new Set([7, 28, 90, days])].sort((a, b) => a - b); + const cost = costPerChange(current); + const baselineCost = costPerChange(baseline); + return ( + + + ROI Calculator +
+ {actions} + {!readOnly && ( + + )} +
+
+ {notice} +
+ + }> + {number(snapshot.repos.length)} {snapshot.repos.length === 1 ? "repository" : "repositories"} + + + + Repositories +
    + {snapshot.repos.map((repo) => ( +
  • {repo}
  • + ))} +
+
+
+
+ + + + {dateRange(current.window)} + + +
+
+
+ + + + +
+ +
+ + + Engineers {snapshot.people.length} + + {terms.requests} + Quality + Branch spend + + {!readOnly && ( + + )} +
+ + setPersonEmail(selected.email)} + /> + + + {current.merged_prs > 0 || baseline.merged_prs > 0 ? ( +
+ +
+
+

Behind the numbers

+

+ {number(current.agent_authored)} of {number(current.merged_prs)} {terms.plural} were authored by + agents or bots. +

+

+ Human-authored median merge time:{" "} + + {duration(current.human_summary.median_merge_hours)} + + , compared with {duration(baseline.human_summary.median_merge_hours)}. +

+
+
+ + {number(current.agents_without_requester)} agent {terms.plural} have no requester + +
+
+
+ ) : ( +
+

No merged changes yet

+

+ Your repositories are connected. New activity will appear after the next sync +

+
+ )} + +

+ All repository {terms.plural}, including agent work without a known requester +

+ +
+ + + + + + +
+
+ + {number(current.matched_internal_prs)} {terms.plural} matched to {snapshot.people.length} engineers ·{" "} + {number(current.agents_without_requester)} agent {terms.plural} without a requester + + Comparing with {dateRange(baseline.window)} · All dates UTC + Spend recorded by this gateway · Merge time is elapsed time, not effort +
+ {accountEmail !== null && ( + setAccountEmail(null)} + onSaved={onRefresh} + /> + )} + {person && ( + setPersonEmail(null)} + onEdit={ + readOnly + ? undefined + : () => { + setAccountEmail(person.email); + setPersonEmail(null); + } + } + /> + )} +
+ ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedData.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedData.test.ts new file mode 100644 index 00000000000..57c37af900a --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedData.test.ts @@ -0,0 +1,110 @@ +import { describe, expect, it } from "vitest"; +import { + change, + duration, + money, + visiblePeople, + weeklyMerges, + type ObservedPerson, + type ObservedSnapshot, +} from "./observedData"; + +const period = (merged: number, spend: number | null) => ({ + merged_prs: merged, + prs_per_week: merged / 4, + median_merge_hours: null, + direct_authored: merged, + declared_agent_owned: 0, + gateway_recorded_spend: spend ?? 0, + recorded_spend_per_attributed_pr: spend !== null && merged > 0 ? spend / merged : null, + spend_observation: spend === null ? ("no_records" as const) : ("records_present" as const), + pr_urls: [], +}); +const person = (name: string, merged: number, spend: number | null): ObservedPerson => ({ + name, + email: `${name}@example.test`, + logins: [`old-${name}`], + periods: { current: period(merged, spend), previous: period(0, null), last_year: period(0, null) }, +}); + +describe("observed ROI metrics", () => { + it("does not claim infinite growth when the baseline is missing", () => { + expect(change(12, 0)).toBeNull(); + expect(change(0, 0)).toBeNull(); + expect(change(15, 10)).toBe(50); + expect(change(0, 10)).toBe(-100); + }); + + it("distinguishes missing cost, measured zero, and small nonzero spend", () => { + expect(money(null)).toBe("Unavailable"); + expect(money(0)).toBe("$0.00"); + expect(money(0.001)).toBe("<$0.01"); + expect(money(1.235)).toBe("$1.24"); + expect(duration(null)).toBe("Unavailable"); + }); + + it("searches historical identities and sorts without changing the source", () => { + const people = [person("Ari", 2, 8), person("Bea", 10, 0), person("Cam", 0, null)]; + expect(visiblePeople(people, " OLD-ARI ", "merged").map((row) => row.name)).toEqual(["Ari"]); + expect(visiblePeople(people, "", "merged").map((row) => row.name)).toEqual(["Bea", "Ari", "Cam"]); + expect(visiblePeople(people, "", "cost").map((row) => row.name)).toEqual(["Ari", "Bea", "Cam"]); + expect(people.map((row) => row.name)).toEqual(["Ari", "Bea", "Cam"]); + expect(visiblePeople(people, "missing", "name")).toEqual([]); + }); + + it("keeps short elapsed merge times from rounding to zero hours", () => { + expect(duration(null)).toBe("Unavailable"); + expect(duration(0)).toBe("0m"); + expect(duration(16 / 3600)).toBe("<1m"); + expect(duration(59 / 3600)).toBe("<1m"); + expect(duration(1 / 60)).toBe("1m"); + expect(duration(79 / 3600)).toBe("1.3m"); + expect(duration(140 / 3600)).toBe("2.3m"); + expect(duration(0.5)).toBe("30m"); + expect(duration(1)).toBe("1h"); + expect(duration(3.82)).toBe("3.8h"); + }); + + it.each([7, 28, 90])("includes every day of a %s-day range in its weekly chart", (days) => { + const start = Date.parse("2026-01-01T00:00:00Z"); + const atDay = (day: number) => new Date(start + day * 86400000).toISOString(); + const pulls = Array.from({ length: days + 1 }, (_, day) => ({ merged_at: atDay(day) })); + const snapshot = { + periods: { current: { window: { start: "2026-01-01", end: atDay(days - 1).slice(0, 10) } } }, + pulls: { current: pulls }, + } as ObservedSnapshot; + const weeks = weeklyMerges(snapshot, "current"); + expect(weeks).toHaveLength(Math.ceil(days / 7)); + expect(weeks.reduce((sum, count) => sum + count, 0)).toBe(days); + expect(weeks.at(-1)).toBe(days % 7 || 7); + }); + + it("aligns comparisons to their own UTC windows and counts each boundary once", () => { + const currentStart = "2026-01-01"; + const previousStart = "2025-12-04"; + const pulls = [ + "2026-01-01T00:00:00Z", + "2026-01-07T23:59:59Z", + "2026-01-08T00:00:00Z", + "2026-01-28T23:59:59Z", + "2026-01-29T00:00:00Z", + ].map((merged_at, number) => ({ + number, + merged_at, + title: "Example", + url: "https://github.com/example/repo/pull/1", + author: "ari", + agent: false, + merge_hours: 1, + })); + const snapshot = { + periods: { + current: { window: { start: currentStart, end: "2026-01-28" } }, + previous: { window: { start: previousStart, end: "2025-12-31" } }, + }, + pulls: { current: pulls, previous: [{ ...pulls[0], merged_at: `${previousStart}T00:00:00Z` }] }, + } as ObservedSnapshot; + expect(weeklyMerges(snapshot, "current")).toEqual([2, 1, 0, 1]); + expect(weeklyMerges(snapshot, "previous")).toEqual([1, 0, 0, 0]); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedData.ts b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedData.ts new file mode 100644 index 00000000000..3dbd115421c --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedData.ts @@ -0,0 +1,212 @@ +import { z } from "zod"; + +const windowSchema = z.object({ start: z.string(), end: z.string() }); +const personPeriodFields = { + merged_prs: z.number(), + prs_per_week: z.number(), + median_merge_hours: z.number().nullable(), + direct_authored: z.number(), + declared_agent_owned: z.number(), + gateway_recorded_spend: z.number(), + recorded_spend_per_attributed_pr: z.number().nullable(), + spend_observation: z.enum(["records_present", "no_records"]), + pr_urls: z.array(z.string()), +}; +const personPeriodSchema = z.object(personPeriodFields); +const periodFields = { + window: windowSchema, + merged_prs: z.number(), + median_merge_hours: z.number().nullable(), + human_authored: z.number(), + agent_authored: z.number(), + missing_author: z.number(), + agents_without_requester: z.number(), + matched_internal_prs: z.number(), + new_bug_labeled_issues: z.number().nullable(), + new_regression_labeled_issues: z.number().nullable(), + explicitly_titled_revert_prs: z.number(), + matched_users_recorded_spend: z.number(), + spend_observation: z.enum(["records_present", "no_records"]), + human_summary: z.object({ median_merge_hours: z.number().nullable() }), +}; +const periodSchema = z.object(periodFields); +const periods = (schema: T) => z.object({ current: schema, previous: schema, last_year: schema }); + +const personFields = { + name: z.string(), + email: z.string(), + logins: z.array(z.string()), + periods: periods(personPeriodSchema), + accounts: z.array(z.object({ connection_id: z.string(), login: z.string() })).optional(), +}; +const branchCostFields = { + repo: z.string(), + branch: z.string(), + spend: z.number().nullable(), + requests: z.number(), + status: z.enum(["matched", "unattributed", "ambiguous", "unavailable"]), +}; +const branchSpendFields = { repo: z.string(), branch: z.string(), spend: z.number(), requests: z.number() }; +const pullFields = { + connection_id: z.string().optional(), + number: z.number(), + title: z.string(), + url: z + .string() + .url() + .refine((url) => new URL(url).protocol === "https:"), + author: z.string(), + agent: z.boolean(), + merged_at: z.string(), + merge_hours: z.number().nullable(), + repo: z.string(), + source_repo: z.string(), + source_branch: z.string(), + branch_cost: z.object(branchCostFields), +}; +const snapshotFields = { + source_provider: z.enum(["github", "gitlab", "mixed"]), + repos: z.array(z.string()), + unmatched_logins: z.array(z.string()), + unlinked_branches: z.array(z.object(branchSpendFields)), + captured_at: z.string(), + periods: periods(periodSchema), + people: z.array(z.object(personFields)), + pulls: periods(z.array(z.object(pullFields))), +}; +export const observedSnapshotSchema = z.object(snapshotFields); +const settingsFields = { + id: z.string().optional(), + source_provider: z.enum(["github", "gitlab"]), + api_url: z.string(), + repos: z.array(z.string()), + has_token: z.boolean(), + connection_type: z.enum(["token", "app"]), + update_interval_minutes: z.number(), + ready: z.boolean(), +}; +export const observedConnectionSchema = z.object(settingsFields); +export const observedSettingsSchema = z.object({ + ...settingsFields, + connections: z.array(observedConnectionSchema).optional(), +}); +export type ObservedConnection = z.infer; +const statusFields = { + running: z.boolean(), + phase: z.string(), + stage: z.string(), + done: z.number(), + total: z.number(), + error: z.string().nullable(), + finished_at: z.string().nullable().optional(), +}; +export const observedStatusSchema = z.object(statusFields); +export const observedReportResponseSchema = z.object({ report: observedSnapshotSchema.nullable() }); +export type ObservedSettings = z.infer; +export type ObservedStatus = z.infer; + +export type ObservedSnapshot = z.infer; +export type ObservedPerson = ObservedSnapshot["people"][number]; +export type ObservedPull = ObservedSnapshot["pulls"]["current"][number]; +export type Period = keyof ObservedSnapshot["periods"]; +export type Comparison = Exclude; +export type PeopleSort = "merged" | "spend" | "cost" | "name"; + +export const number = (value: number | null) => + value === null ? "Unavailable" : value.toLocaleString("en-US", { maximumFractionDigits: 1 }); +export function money(value: number | null) { + if (value === null) return "Unavailable"; + if (value > 0 && value < 0.01) return "<$0.01"; + return value.toLocaleString("en-US", { style: "currency", currency: "USD", maximumFractionDigits: 2 }); +} +export function duration(value: number | null) { + if (value === null) return "Unavailable"; + if (value === 0) return "0m"; + if (value < 1 / 60) return "<1m"; + if (value < 1) return `${number(value * 60)}m`; + return `${number(value)}h`; +} +export const change = (current: number, baseline: number) => + baseline === 0 ? null : ((current - baseline) / baseline) * 100; + +export const accountLogins = (value: string) => [ + ...new Set( + value + .split(/[\s,]+/) + .map((login) => login.replace(/^@/, "").toLowerCase()) + .filter(Boolean), + ), +]; +export const repositoryNames = (value: string) => [ + ...new Set( + value + .split(/[\s,]+/) + .filter(Boolean) + .map((repo) => { + const path = repo.replace(/^https:\/\/[^/]+\//, ""); + return path.replace(/\/$/, "").replace(/\.git$/, ""); + }), + ), +]; + +export function dateRange(window: { start: string; end: string }) { + const date = (value: string) => + new Date(`${value}T00:00:00Z`).toLocaleDateString("en-US", { + month: "short", + day: "numeric", + timeZone: "UTC", + }); + return `${date(window.start)} – ${date(window.end)}, ${window.end.slice(0, 4)}`; +} + +export function visiblePeople(people: ObservedPerson[], query: string, sort: PeopleSort) { + const value = (person: ObservedPerson) => { + const current = person.periods.current; + if (sort === "spend") return current.spend_observation === "no_records" ? -1 : current.gateway_recorded_spend; + if (sort === "cost") return current.recorded_spend_per_attributed_pr ?? -1; + return current.merged_prs; + }; + return people + .filter((person) => + [person.name, person.email, ...person.logins].join(" ").toLowerCase().includes(query.trim().toLowerCase()), + ) + .toSorted((a, b) => (sort === "name" ? a.name.localeCompare(b.name) : value(b) - value(a))); +} + +export function weeklyMerges(snapshot: ObservedSnapshot, period: Period) { + const start = Date.parse(`${snapshot.periods[period].window.start}T00:00:00Z`); + const end = Date.parse(`${snapshot.periods[period].window.end}T00:00:00Z`) + 86_400_000; + const days = (end - start) / 86_400_000; + return Array.from( + { length: Math.ceil(days / 7) }, + (_, week) => + snapshot.pulls[period].filter((pull) => { + const day = (Date.parse(pull.merged_at) - start) / 86_400_000; + return day >= week * 7 && day < Math.min((week + 1) * 7, days); + }).length, + ); +} + +export function syncMessage(status: ObservedStatus, report: ObservedSnapshot | null) { + if (status.running) return status.total ? `${status.stage} · ${status.done} / ${status.total}` : status.stage; + return report ? `Updated ${new Date(report.captured_at).toLocaleString()}` : ""; +} + +export function recordedBranches(snapshot: ObservedSnapshot) { + const matched = snapshot.pulls.current.flatMap((pull) => { + const cost = pull.branch_cost; + if (cost.status !== "matched" || cost.spend === null) return []; + return [{ repo: cost.repo, branch: cost.branch, spend: cost.spend, requests: cost.requests }]; + }); + return [ + ...new Map([...matched, ...snapshot.unlinked_branches].map((row) => [`${row.repo}\n${row.branch}`, row])).values(), + ].toSorted((a, b) => b.spend - a.spend); +} + +export function changeTerms(provider: ObservedSnapshot["source_provider"]) { + if (provider === "mixed") + return { singular: "change", plural: "changes", requests: "Merged changes", lower: "merged changes" }; + return provider === "gitlab" + ? { singular: "MR", plural: "MRs", requests: "Merge requests", lower: "merge requests" } + : { singular: "PR", plural: "PRs", requests: "Pull requests", lower: "pull requests" }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedDemo.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedDemo.test.ts new file mode 100644 index 00000000000..23ad2dc85b6 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedDemo.test.ts @@ -0,0 +1,42 @@ +import { describe, expect, it } from "vitest"; +import { observedSnapshotSchema, weeklyMerges, type Period } from "./observedData"; +import { createObservedDemo } from "./observedDemo"; + +describe("observed sample report", () => { + it("clamps the year-ago comparison on leap day while keeping the selected length", () => { + const report = createObservedDemo(28, new Date("2024-03-01T12:00:00Z")); + expect(report.periods.current.window.end).toBe("2024-02-29"); + expect(report.periods.last_year.window).toEqual({ start: "2023-02-01", end: "2023-02-28" }); + }); + + it.each([7, 28, 90])("keeps totals, attribution, costs, and comparison windows consistent for %i days", (days) => { + const report = createObservedDemo(days, new Date("2026-10-03T14:00:00Z")); + expect(observedSnapshotSchema.safeParse(report).success).toBe(true); + for (const period of ["current", "previous", "last_year"] satisfies Period[]) { + const metrics = report.periods[period]; + const pulls = report.pulls[period]; + const people = report.people.map((person) => person.periods[period]); + const start = Date.parse(metrics.window.start); + const end = Date.parse(metrics.window.end) + 86_400_000; + expect((end - start) / 86_400_000).toBe(days); + expect(metrics.merged_prs).toBe(pulls.length); + expect(new Set(pulls.map((pull) => pull.url)).size).toBe(pulls.length); + expect(weeklyMerges(report, period).reduce((sum, value) => sum + value, 0)).toBe(pulls.length); + expect(pulls.every((pull) => Date.parse(pull.merged_at) >= start && Date.parse(pull.merged_at) < end)).toBe(true); + expect(metrics.matched_users_recorded_spend).toBe( + people.reduce((sum, person) => sum + person.gateway_recorded_spend, 0), + ); + expect(people.reduce((sum, person) => sum + person.merged_prs, 0)).toBe(metrics.matched_internal_prs); + for (const person of people) { + const attributed = pulls.filter((pull) => person.pr_urls.includes(pull.url)); + expect(attributed).toHaveLength(person.merged_prs); + expect(person.direct_authored).toBe(attributed.filter((pull) => !pull.agent).length); + expect(person.declared_agent_owned).toBe(attributed.filter((pull) => pull.agent).length); + expect(person.recorded_spend_per_attributed_pr).toBe(person.gateway_recorded_spend / person.merged_prs); + } + } + expect(Date.parse(report.periods.previous.window.end) + 86_400_000).toBe( + Date.parse(report.periods.current.window.start), + ); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedDemo.ts b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedDemo.ts new file mode 100644 index 00000000000..fb43fa8df3f --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/observedDemo.ts @@ -0,0 +1,120 @@ +import type { ObservedPerson, ObservedPull, ObservedSnapshot, Period } from "./observedData"; + +const DAY = 86_400_000; +const engineers = [ + { name: "Alex Rivera", login: "alex-demo", weekly: [8, 6, 4] }, + { name: "Sam Chen", login: "sam-demo", weekly: [6, 5, 4] }, + { name: "Jordan Lee", login: "jordan-demo", weekly: [4, 4, 3] }, +]; +const titles = [ + "Add repository search", + "Fix retry handling", + "Speed up activity queries", + "Add usage export", + "Improve connection setup", + "Fix pagination", +]; +const isoDate = (timestamp: number) => new Date(timestamp).toISOString().slice(0, 10); + +function median(pulls: ObservedPull[]) { + const hours = pulls.map((pull) => pull.merge_hours ?? 0).toSorted((a, b) => a - b); + const middle = Math.floor(hours.length / 2); + if (!hours.length) return null; + return hours.length % 2 ? hours[middle] : (hours[middle - 1] + hours[middle]) / 2; +} + +function samplePeriod(days: number, end: number, comparison: number) { + const start = end - (days - 1) * DAY; + const people = engineers.map((engineer, index) => { + const count = Math.round((engineer.weekly[comparison] * days) / 7); + const pulls: ObservedPull[] = Array.from({ length: count }, (_, position) => { + const number = 10000 * (comparison + 1) + index * 1000 + position; + const repo = index === 1 ? "demo/api" : "demo/web"; + const branch = `feature/sample-${number}`; + const agent = position % 5 === 0; + return { + connection_id: index === 1 ? "demo-gitlab" : "demo-github", + number, + title: titles[position % titles.length], + url: `https://example.com/${repo}/changes/${number}`, + author: agent ? "demo-agent" : engineer.login, + agent, + merged_at: new Date(start + Math.floor((position * days) / count) * DAY + 12 * 3_600_000).toISOString(), + merge_hours: position === 0 ? 16 / 3600 : 4 + ((position * 7) % 24) + comparison * 6, + repo, + source_repo: repo, + source_branch: branch, + branch_cost: { repo, branch, spend: 2 + (position % 4), requests: 20 + position, status: "matched" }, + }; + }); + const spend = pulls.reduce((total, pull) => total + (pull.branch_cost.spend ?? 0), 0); + const metrics: ObservedPerson["periods"]["current"] = { + merged_prs: pulls.length, + prs_per_week: (pulls.length * 7) / days, + median_merge_hours: median(pulls), + direct_authored: pulls.filter((pull) => !pull.agent).length, + declared_agent_owned: pulls.filter((pull) => pull.agent).length, + gateway_recorded_spend: spend, + recorded_spend_per_attributed_pr: pulls.length ? spend / pulls.length : null, + spend_observation: "records_present", + pr_urls: pulls.map((pull) => pull.url), + }; + return { metrics, pulls }; + }); + const pulls = people.flatMap((person) => person.pulls).toSorted((a, b) => b.merged_at.localeCompare(a.merged_at)); + const metrics: ObservedSnapshot["periods"]["current"] = { + window: { start: isoDate(start), end: isoDate(end) }, + merged_prs: pulls.length, + median_merge_hours: median(pulls), + human_authored: pulls.filter((pull) => !pull.agent).length, + agent_authored: pulls.filter((pull) => pull.agent).length, + missing_author: 0, + agents_without_requester: 0, + matched_internal_prs: pulls.length, + new_bug_labeled_issues: Math.round(((comparison + 1) * days) / 7), + new_regression_labeled_issues: comparison, + explicitly_titled_revert_prs: 0, + matched_users_recorded_spend: people.reduce((total, person) => total + person.metrics.gateway_recorded_spend, 0), + spend_observation: "records_present", + human_summary: { median_merge_hours: median(pulls.filter((pull) => !pull.agent)) }, + }; + return { metrics, pulls, people }; +} + +export function createObservedDemo(days: number, now = new Date()): ObservedSnapshot { + const end = Date.UTC(now.getUTCFullYear(), now.getUTCMonth(), now.getUTCDate()) - DAY; + const yearAgo = new Date(end); + const lastYear = yearAgo.getUTCFullYear() - 1; + const month = yearAgo.getUTCMonth(); + const lastDay = new Date(Date.UTC(lastYear, month + 1, 0)).getUTCDate(); + const lastYearEnd = Date.UTC(lastYear, month, Math.min(yearAgo.getUTCDate(), lastDay)); + const periods = { + current: samplePeriod(days, end, 0), + previous: samplePeriod(days, end - days * DAY, 1), + last_year: samplePeriod(days, lastYearEnd, 2), + }; + const personPeriod = (period: Period, index: number) => periods[period].people[index].metrics; + return { + source_provider: "mixed", + repos: ["demo/web", "demo/api"], + unmatched_logins: [], + unlinked_branches: [], + captured_at: now.toISOString(), + periods: { + current: periods.current.metrics, + previous: periods.previous.metrics, + last_year: periods.last_year.metrics, + }, + people: engineers.map((engineer, index) => ({ + name: engineer.name, + email: `${engineer.login}@example.com`, + logins: [engineer.login], + periods: { + current: personPeriod("current", index), + previous: personPeriod("previous", index), + last_year: personPeriod("last_year", index), + }, + })), + pulls: { current: periods.current.pulls, previous: periods.previous.pulls, last_year: periods.last_year.pulls }, + }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/useObservedReport.ts b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/useObservedReport.ts new file mode 100644 index 00000000000..989946e2a27 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/useObservedReport.ts @@ -0,0 +1,60 @@ +import { useCallback, useEffect, useState } from "react"; +import { apiClient } from "@/components/networking"; +import { extractProxyErrorMessage } from "@/lib/http/client"; +import { + observedSettingsSchema, + observedReportResponseSchema, + observedStatusSchema, + type ObservedSettings, + type ObservedSnapshot, + type ObservedStatus, +} from "./observedData"; + +export type ObservedViewData = { settings: ObservedSettings; report: ObservedSnapshot | null; status: ObservedStatus }; + +export function useObservedReport(accessToken: string) { + const [data, setData] = useState(null); + const [error, setError] = useState(""); + const [revision, setRevision] = useState(0); + const refresh = useCallback(() => setRevision((value) => value + 1), []); + useEffect(() => { + const controller = new AbortController(); + let timer: ReturnType; + const options = { accessToken, signal: controller.signal }; + async function poll(previous?: ObservedViewData) { + try { + const status = observedStatusSchema.parse( + await apiClient.get("/roi-calculator/observed/sync", options), + ); + const finished = previous?.status.running && !status.running; + const changed = !previous || previous.status.finished_at !== status.finished_at || finished; + const updated = changed + ? await Promise.all([ + apiClient + .get("/roi-calculator/observed/settings", options) + .then((value) => observedSettingsSchema.parse(value)), + apiClient + .get("/roi-calculator/observed/report", options) + .then((value) => observedReportResponseSchema.parse(value)), + ]) + : null; + if (controller.signal.aborted) return; + const existing = previous ? { ...previous, status } : null; + const next = updated ? { settings: updated[0], report: updated[1].report, status } : existing; + setData(next); + setError(""); + timer = setTimeout(() => void poll(next ?? undefined), status.running ? 2000 : 30000); + } catch (reason) { + if (controller.signal.aborted) return; + setError(extractProxyErrorMessage(reason)); + timer = setTimeout(() => void poll(previous), 10000); + } + } + void poll(); + return () => { + controller.abort(); + clearTimeout(timer); + }; + }, [accessToken, revision]); + return { data, error, refresh }; +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/page.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/page.tsx index 329ecbc0fe6..da1688e342f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/page.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/page.tsx @@ -1,9 +1,10 @@ "use client"; import useAuthorized from "@/app/(dashboard)/hooks/useAuthorized"; -import ROICalculatorView from "./_components/ROICalculatorView"; +import ObservedROIView from "./_components/ObservedROIView"; export default function ROICalculatorPage() { - const { accessToken, userRole, isViewOnly } = useAuthorized(); - return ; + const { accessToken, isViewOnly } = useAuthorized(); + if (!accessToken) return null; + return ; } diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 0dde64137a5..5478772bc29 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -14331,6 +14331,129 @@ export interface paths { patch?: never; trace?: never; }; + "/roi-calculator/observed/apps": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Observed Apps */ + get: operations["observed_apps_roi_calculator_observed_apps_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/roi-calculator/observed/identities": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Observed Identities */ + get: operations["get_observed_identities_roi_calculator_observed_identities_get"]; + /** Save Observed Identities */ + put: operations["save_observed_identities_roi_calculator_observed_identities_put"]; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/roi-calculator/observed/oauth/{provider}/start": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + get?: never; + put?: never; + /** Start Observed Authorization */ + post: operations["start_observed_authorization_roi_calculator_observed_oauth__provider__start_post"]; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/roi-calculator/observed/report": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Observed Report */ + get: operations["get_observed_report_roi_calculator_observed_report_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/roi-calculator/observed/repositories": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Observed Repositories */ + get: operations["observed_repositories_roi_calculator_observed_repositories_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/roi-calculator/observed/settings": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Observed Settings */ + get: operations["get_observed_settings_roi_calculator_observed_settings_get"]; + /** Save Observed Settings */ + put: operations["save_observed_settings_roi_calculator_observed_settings_put"]; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; + "/roi-calculator/observed/sync": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** Get Observed Sync */ + get: operations["get_observed_sync_roi_calculator_observed_sync_get"]; + put?: never; + /** Start Observed Sync */ + post: operations["start_observed_sync_roi_calculator_observed_sync_post"]; + /** Cancel Observed Sync */ + delete: operations["cancel_observed_sync_roi_calculator_observed_sync_delete"]; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/roi-calculator/report": { parameters: { query?: never; @@ -39236,6 +39359,373 @@ export interface components { [key: string]: unknown; } | null; }; + /** ObservedAccount */ + ObservedAccount: { + /** Connection Id */ + connection_id: string; + /** Login */ + login: string; + }; + /** ObservedApp */ + ObservedApp: { + /** Api Url */ + api_url?: string | null; + /** Callback Url */ + callback_url?: string | null; + /** + * Can Install + * @default false + */ + can_install: boolean; + /** Configured */ + configured: boolean; + }; + /** ObservedApps */ + ObservedApps: { + github: components["schemas"]["ObservedApp"]; + gitlab: components["schemas"]["ObservedApp"]; + }; + /** ObservedAuthorization */ + ObservedAuthorization: { + /** Url */ + url: string; + }; + /** ObservedConnection */ + ObservedConnection: { + /** Api Url */ + api_url: string; + /** + * Connection Type + * @enum {string} + */ + connection_type: "token" | "app"; + /** Has Token */ + has_token: boolean; + /** + * Id + * @default + */ + id: string; + /** Ready */ + ready: boolean; + /** Repos */ + repos: string[]; + /** + * Source Provider + * @enum {string} + */ + source_provider: "github" | "gitlab"; + /** Update Interval Minutes */ + update_interval_minutes: number; + }; + /** ObservedConnectionIdentities */ + ObservedConnectionIdentities: { + /** Api Url */ + api_url: string; + /** Id */ + id: string; + /** Identity Map */ + identity_map: { + [key: string]: string; + }; + /** Repos */ + repos: string[]; + /** + * Source Provider + * @enum {string} + */ + source_provider: "github" | "gitlab"; + /** Unmatched Logins */ + unmatched_logins: string[]; + }; + /** ObservedHumanSummary */ + ObservedHumanSummary: { + /** Median Merge Hours */ + median_merge_hours: number | null; + }; + /** ObservedIdentities */ + ObservedIdentities: { + /** + * Connections + * @default [] + */ + connections: components["schemas"]["ObservedConnectionIdentities"][]; + /** Gateway Emails */ + gateway_emails: string[]; + /** Identity Map */ + identity_map: { + [key: string]: string; + }; + /** Unmatched Logins */ + unmatched_logins: string[]; + }; + /** ObservedIdentityUpdate */ + ObservedIdentityUpdate: { + /** Accounts */ + accounts?: components["schemas"]["ObservedAccount"][] | null; + /** Email */ + email: string; + /** + * Logins + * @default [] + */ + logins: string[]; + }; + /** ObservedPeriod */ + ObservedPeriod: { + /** Agent Authored */ + agent_authored: number; + /** Agents Without Requester */ + agents_without_requester: number; + /** Explicitly Titled Revert Prs */ + explicitly_titled_revert_prs: number; + /** Human Authored */ + human_authored: number; + human_summary: components["schemas"]["ObservedHumanSummary"]; + /** Matched Internal Prs */ + matched_internal_prs: number; + /** Matched Users Recorded Spend */ + matched_users_recorded_spend: number; + /** Median Merge Hours */ + median_merge_hours: number | null; + /** Merged Prs */ + merged_prs: number; + /** Missing Author */ + missing_author: number; + /** New Bug Labeled Issues */ + new_bug_labeled_issues: number | null; + /** New Regression Labeled Issues */ + new_regression_labeled_issues: number | null; + /** + * Spend Observation + * @enum {string} + */ + spend_observation: "records_present" | "no_records"; + window: components["schemas"]["ObservedWindow"]; + }; + /** ObservedPeriods */ + ObservedPeriods: { + current: components["schemas"]["ObservedPeriod"]; + last_year: components["schemas"]["ObservedPeriod"]; + previous: components["schemas"]["ObservedPeriod"]; + }; + /** ObservedPerson */ + ObservedPerson: { + /** + * Accounts + * @default [] + */ + accounts: components["schemas"]["ObservedAccount"][]; + /** Email */ + email: string; + /** Logins */ + logins: string[]; + /** Name */ + name: string; + periods: components["schemas"]["ObservedPersonPeriods"]; + }; + /** ObservedPersonPeriod */ + ObservedPersonPeriod: { + /** Declared Agent Owned */ + declared_agent_owned: number; + /** Direct Authored */ + direct_authored: number; + /** Gateway Recorded Spend */ + gateway_recorded_spend: number; + /** Median Merge Hours */ + median_merge_hours: number | null; + /** Merged Prs */ + merged_prs: number; + /** Pr Urls */ + pr_urls: string[]; + /** Prs Per Week */ + prs_per_week: number; + /** Recorded Spend Per Attributed Pr */ + recorded_spend_per_attributed_pr: number | null; + /** + * Spend Observation + * @enum {string} + */ + spend_observation: "records_present" | "no_records"; + }; + /** ObservedPersonPeriods */ + ObservedPersonPeriods: { + current: components["schemas"]["ObservedPersonPeriod"]; + last_year: components["schemas"]["ObservedPersonPeriod"]; + previous: components["schemas"]["ObservedPersonPeriod"]; + }; + /** ObservedPullPeriods */ + ObservedPullPeriods: { + /** Current */ + current: components["schemas"]["ObservedPullResponse"][]; + /** Last Year */ + last_year: components["schemas"]["ObservedPullResponse"][]; + /** Previous */ + previous: components["schemas"]["ObservedPullResponse"][]; + }; + /** ObservedPullResponse */ + ObservedPullResponse: { + /** + * Agent + * @default false + */ + agent: boolean; + /** Author */ + author: string; + branch_cost: components["schemas"]["ROIBranchAttribution"]; + /** + * Connection Id + * @default + */ + connection_id: string; + /** Created At */ + created_at?: string | null; + /** Merge Hours */ + merge_hours: number | null; + /** + * Merged At + * Format: date-time + */ + merged_at: string; + /** Number */ + number: number; + /** + * Profile Email + * @default + */ + profile_email: string; + /** Repo */ + repo: string; + /** + * Requester + * @default + */ + requester: string; + /** + * Source Branch + * @default + */ + source_branch: string; + /** + * Source Repo + * @default + */ + source_repo: string; + /** Title */ + title: string; + /** Url */ + url: string; + }; + /** ObservedReport */ + ObservedReport: { + /** + * Captured At + * Format: date-time + */ + captured_at: string; + /** + * Connections + * @default [] + */ + connections: components["schemas"]["ObservedSource"][]; + /** People */ + people: components["schemas"]["ObservedPerson"][]; + periods: components["schemas"]["ObservedPeriods"]; + pulls: components["schemas"]["ObservedPullPeriods"]; + /** Repos */ + repos: string[]; + /** + * Source Provider + * @enum {string} + */ + source_provider: "github" | "gitlab" | "mixed"; + /** Unlinked Branches */ + unlinked_branches: components["schemas"]["ROIBranchSpend"][]; + /** Unmatched Logins */ + unmatched_logins: string[]; + }; + /** ObservedReportResponse */ + ObservedReportResponse: { + report: components["schemas"]["ObservedReport"] | null; + }; + /** ObservedSettings */ + ObservedSettings: { + /** Api Url */ + api_url: string; + /** + * Connection Type + * @enum {string} + */ + connection_type: "token" | "app"; + /** + * Connections + * @default [] + */ + connections: components["schemas"]["ObservedConnection"][]; + /** Has Token */ + has_token: boolean; + /** + * Id + * @default + */ + id: string; + /** Ready */ + ready: boolean; + /** Repos */ + repos: string[]; + /** + * Source Provider + * @enum {string} + */ + source_provider: "github" | "gitlab"; + /** Update Interval Minutes */ + update_interval_minutes: number; + }; + /** ObservedSettingsUpdate */ + ObservedSettingsUpdate: { + /** Api Url */ + api_url: string; + /** Connection Id */ + connection_id?: string | null; + /** Repos */ + repos: string[]; + /** + * Source Provider + * @enum {string} + */ + source_provider: "github" | "gitlab"; + /** Token */ + token?: string | null; + /** Update Interval Minutes */ + update_interval_minutes?: number | null; + }; + /** ObservedSource */ + ObservedSource: { + /** Api Url */ + api_url: string; + /** Id */ + id: string; + /** Repos */ + repos: string[]; + /** + * Source Provider + * @enum {string} + */ + source_provider: "github" | "gitlab"; + }; + /** ObservedWindow */ + ObservedWindow: { + /** + * End + * Format: date + */ + end: string; + /** + * Start + * Format: date + */ + start: string; + }; /** * OpenIdConnectSecurityScheme * @description Defines a security scheme using OpenID Connect. @@ -41591,6 +42081,12 @@ export interface components { }; /** Ready */ ready: boolean; + /** + * Report Mode + * @default legacy + * @enum {string} + */ + report_mode: "legacy" | "observed"; /** Repos */ repos: string[]; /** @@ -41620,6 +42116,8 @@ export interface components { gitlab_api_url?: string | null; /** Gitlab Token */ gitlab_token?: string | null; + /** Report Mode */ + report_mode?: ("legacy" | "observed") | null; /** Repos */ repos?: string[] | null; /** Source Provider */ @@ -69452,6 +69950,289 @@ export interface operations { }; }; }; + observed_apps_roi_calculator_observed_apps_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ObservedApps"]; + }; + }; + }; + }; + get_observed_identities_roi_calculator_observed_identities_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ObservedIdentities"]; + }; + }; + }; + }; + save_observed_identities_roi_calculator_observed_identities_put: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody: { + content: { + "application/json": components["schemas"]["ObservedIdentityUpdate"]; + }; + }; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ObservedReportResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + start_observed_authorization_roi_calculator_observed_oauth__provider__start_post: { + parameters: { + query?: { + install?: boolean; + }; + header?: never; + path: { + provider: "github" | "gitlab"; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ObservedAuthorization"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + get_observed_report_roi_calculator_observed_report_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ObservedReportResponse"]; + }; + }; + }; + }; + observed_repositories_roi_calculator_observed_repositories_get: { + parameters: { + query?: { + connection?: string | null; + query?: string; + page?: number; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ROIRepositoriesResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + get_observed_settings_roi_calculator_observed_settings_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ObservedSettings"]; + }; + }; + }; + }; + save_observed_settings_roi_calculator_observed_settings_put: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody: { + content: { + "application/json": components["schemas"]["ObservedSettingsUpdate"]; + }; + }; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ObservedSettings"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + get_observed_sync_roi_calculator_observed_sync_get: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ROISyncStatus"]; + }; + }; + }; + }; + start_observed_sync_roi_calculator_observed_sync_post: { + parameters: { + query?: { + days?: number | null; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 202: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ROISyncStatus"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; + cancel_observed_sync_roi_calculator_observed_sync_delete: { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["ROISyncStatus"]; + }; + }; + }; + }; get_roi_calculator_report_roi_calculator_report_get: { parameters: { query?: { From 5f969982fc5f22dc9e67ef8a8752b5231ebe3f80 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 16:15:02 -0700 Subject: [PATCH 02/18] fix(azure): set text-embedding max input to 8192 from the models sold directly page (#44449) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- ...model_prices_and_context_window_backup.json | 18 +++++++++--------- model_prices_and_context_window.json | 18 +++++++++--------- 2 files changed, 18 insertions(+), 18 deletions(-) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index a50edb0c9e3..27495443a57 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -10841,31 +10841,31 @@ "deprecation_date": "2028-02-09", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure", - "max_input_tokens": 8191, - "max_tokens": 8191, + "max_input_tokens": 8192, + "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0.0, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure" }, "azure/text-embedding-3-small": { "deprecation_date": "2028-02-09", "input_cost_per_token": 2e-08, "litellm_provider": "azure", - "max_input_tokens": 8191, - "max_tokens": 8191, + "max_input_tokens": 8192, + "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0.0, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure" }, "azure/text-embedding-ada-002": { "deprecation_date": "2028-02-09", "input_cost_per_token": 1e-07, "litellm_provider": "azure", - "max_input_tokens": 8191, - "max_tokens": 8191, + "max_input_tokens": 8192, + "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0.0, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure" }, "azure/speech/azure-tts": { "input_cost_per_character": 1.5e-05, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index a50edb0c9e3..27495443a57 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -10841,31 +10841,31 @@ "deprecation_date": "2028-02-09", "input_cost_per_token": 1.3e-07, "litellm_provider": "azure", - "max_input_tokens": 8191, - "max_tokens": 8191, + "max_input_tokens": 8192, + "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0.0, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure" }, "azure/text-embedding-3-small": { "deprecation_date": "2028-02-09", "input_cost_per_token": 2e-08, "litellm_provider": "azure", - "max_input_tokens": 8191, - "max_tokens": 8191, + "max_input_tokens": 8192, + "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0.0, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure" }, "azure/text-embedding-ada-002": { "deprecation_date": "2028-02-09", "input_cost_per_token": 1e-07, "litellm_provider": "azure", - "max_input_tokens": 8191, - "max_tokens": 8191, + "max_input_tokens": 8192, + "max_tokens": 8192, "mode": "embedding", "output_cost_per_token": 0.0, - "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" + "source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure" }, "azure/speech/azure-tts": { "input_cost_per_character": 1.5e-05, From 427158eb5ba6e3a586253eb68e8790c16f5cd340 Mon Sep 17 00:00:00 2001 From: yuneng-jiang Date: Sat, 3 Oct 2026 16:23:13 -0700 Subject: [PATCH 03/18] feat(ui): show invitation and reset password links in a copyable field (#44454) * feat(ui): show invitation and reset password links in a copyable field Put the link in a read-only input with a Copy button beside it, stack the User ID and link labels above their values, and focus Copy on open so the field shows the start of the URL. Copy now goes through the shared copyToClipboard helper, which falls back to a selection copy where the Clipboard API is unavailable. * fix(ui): keep focus on the copy control after a fallback clipboard copy The execCommand fallback focused a temporary textarea and removed it, so focus fell to the page body and a second Enter on Copy did nothing. Restore focus to the element that had it. Move the rendered dialog tests to the integration tier. --- .../onboarding_link.integration.test.tsx | 125 ++++++++++++++++++ .../src/components/onboarding_link.tsx | 85 +++++++----- ui/litellm-dashboard/src/utils/dataUtils.ts | 2 + 3 files changed, 176 insertions(+), 36 deletions(-) create mode 100644 ui/litellm-dashboard/src/components/onboarding_link.integration.test.tsx diff --git a/ui/litellm-dashboard/src/components/onboarding_link.integration.test.tsx b/ui/litellm-dashboard/src/components/onboarding_link.integration.test.tsx new file mode 100644 index 00000000000..523aebf0a5c --- /dev/null +++ b/ui/litellm-dashboard/src/components/onboarding_link.integration.test.tsx @@ -0,0 +1,125 @@ +import { afterEach, describe, it, expect, vi } from "vitest"; +import { fireEvent, render, screen, waitFor } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import OnboardingModal, { InvitationLink } from "./onboarding_link"; + +const invitation: InvitationLink = { + id: "inv-123", + user_id: "user-abc", + is_accepted: false, + accepted_at: null, + expires_at: new Date("2030-01-01"), + created_at: new Date("2029-12-01"), + created_by: "admin", + updated_at: new Date("2029-12-01"), + updated_by: "admin", + has_user_setup_sso: false, +}; + +const renderModal = (modalType: "invitation" | "resetPassword", setVisible = vi.fn()) => + render( + , + ); + +describe("OnboardingModal", () => { + it("shows the reset password link in a read-only field labelled for that flow", () => { + renderModal("resetPassword"); + + const field = screen.getByRole("textbox", { name: "Reset password link" }); + expect(field).toHaveValue("http://localhost:4000/ui/onboarding?invitation_id=inv-123&action=reset_password"); + expect(field).toHaveAttribute("readonly"); + expect(screen.getByText("user-abc")).toBeInTheDocument(); + }); + + it("shows the invitation link in a read-only field labelled for that flow", () => { + renderModal("invitation"); + + const field = screen.getByRole("textbox", { name: "Invitation link" }); + expect(field).toHaveValue("http://localhost:4000/ui/onboarding?invitation_id=inv-123"); + expect(field).toHaveAttribute("readonly"); + }); + + it("focuses the copy button on open so the link field shows the start of the URL", async () => { + renderModal("resetPassword"); + + const copyButton = screen.getByRole("button", { name: "Copy password reset link" }); + await waitFor(() => expect(copyButton).toHaveFocus()); + }); + + it.each([ + ["invitation", "Copy invitation link", "http://localhost:4000/ui/onboarding?invitation_id=inv-123"], + [ + "resetPassword", + "Copy password reset link", + "http://localhost:4000/ui/onboarding?invitation_id=inv-123&action=reset_password", + ], + ] as const)("copies exactly the displayed %s link when the copy button is pressed", async (modalType, label, url) => { + const user = userEvent.setup(); + const writeText = vi.spyOn(navigator.clipboard, "writeText").mockResolvedValue(); + renderModal(modalType); + + await user.click(screen.getByRole("button", { name: label })); + + expect(writeText).toHaveBeenCalledWith(url); + }); + + describe("without the Clipboard API, as on a plain-http deployment", () => { + const originalClipboard = Object.getOwnPropertyDescriptor(navigator, "clipboard"); + + afterEach(() => { + if (originalClipboard) Object.defineProperty(navigator, "clipboard", originalClipboard); + Reflect.deleteProperty(document, "execCommand"); + vi.restoreAllMocks(); + }); + + it("still copies exactly the displayed link through the selection fallback", () => { + Object.defineProperty(navigator, "clipboard", { value: undefined, configurable: true }); + const selectedTexts: string[] = []; + vi.spyOn(HTMLTextAreaElement.prototype, "select").mockImplementation(function (this: HTMLTextAreaElement) { + selectedTexts.push(this.value); + }); + const execCommand = vi.fn(() => true); + document.execCommand = execCommand; + renderModal("resetPassword"); + + fireEvent.click(screen.getByRole("button", { name: "Copy password reset link" })); + + expect(selectedTexts).toEqual([ + "http://localhost:4000/ui/onboarding?invitation_id=inv-123&action=reset_password", + ]); + expect(execCommand).toHaveBeenCalledWith("copy"); + }); + + it("keeps focus on the copy button after a fallback copy so Enter copies again", async () => { + const user = userEvent.setup(); + Object.defineProperty(navigator, "clipboard", { value: undefined, configurable: true }); + const execCommand = vi.fn(() => true); + document.execCommand = execCommand; + renderModal("resetPassword"); + const copyButton = screen.getByRole("button", { name: "Copy password reset link" }); + await waitFor(() => expect(copyButton).toHaveFocus()); + + await user.keyboard("{Enter}"); + await user.keyboard("{Enter}"); + + expect(copyButton).toHaveFocus(); + expect(execCommand).toHaveBeenCalledTimes(2); + }); + }); + + it("asks the caller to hide the dialog when Escape is pressed", async () => { + const user = userEvent.setup(); + const setVisible = vi.fn(); + renderModal("invitation", setVisible); + + await user.keyboard("{Escape}"); + + expect(setVisible).toHaveBeenCalledWith(false); + }); +}); diff --git a/ui/litellm-dashboard/src/components/onboarding_link.tsx b/ui/litellm-dashboard/src/components/onboarding_link.tsx index 441a85b1748..e0bb7e4da81 100644 --- a/ui/litellm-dashboard/src/components/onboarding_link.tsx +++ b/ui/litellm-dashboard/src/components/onboarding_link.tsx @@ -1,8 +1,9 @@ -import React from "react"; -import { Button } from "@/components/ui/button"; -import { CopyToClipboard } from "react-copy-to-clipboard"; -import { toast } from "@/lib/toast"; -import { Dialog, DialogContent, DialogHeader, DialogTitle } from "@/components/ui/dialog"; +import React, { useId, useRef } from "react"; +import { Copy } from "lucide-react"; +import { Dialog, DialogContent, DialogDescription, DialogHeader, DialogTitle } from "@/components/ui/dialog"; +import { InputGroup, InputGroupAddon, InputGroupButton, InputGroupInput } from "@/components/ui/input-group"; +import { Label } from "@/components/ui/label"; +import { copyToClipboard } from "@/utils/dataUtils"; export interface InvitationLink { id: string; @@ -58,41 +59,53 @@ export default function OnboardingModal({ invitationLinkData, modalType = "invitation", }: OnboardingProps) { - const handleInvitationCancel = () => { - setIsInvitationLinkModalVisible(false); - }; - - const getInvitationUrl = () => - buildOnboardingUrl({ - baseUrl, - invitationId: invitationLinkData?.id, - hasUserSetupSso: invitationLinkData?.has_user_setup_sso ?? false, - resetPassword: modalType === "resetPassword", - }); + const linkFieldId = useId(); + const copyButtonRef = useRef(null); + const isInvitation = modalType === "invitation"; + const invitationUrl = buildOnboardingUrl({ + baseUrl, + invitationId: invitationLinkData?.id, + hasUserSetupSso: invitationLinkData?.has_user_setup_sso ?? false, + resetPassword: !isInvitation, + }); return ( - !open && handleInvitationCancel()}> - + !open && setIsInvitationLinkModalVisible(false)} + > + - {modalType === "invitation" ? "Invitation Link" : "Reset Password Link"} + {isInvitation ? "Invitation Link" : "Reset Password Link"} + + {isInvitation + ? "Copy and send the generated link to onboard this user to the proxy." + : "Copy and send the generated link to the user to reset their password."} + -

- {modalType === "invitation" - ? "Copy and send the generated link to onboard this user to the proxy." - : "Copy and send the generated link to the user to reset their password."} -

-
-

User ID

-

{invitationLinkData?.user_id}

-
-
-

{modalType === "invitation" ? "Invitation Link" : "Reset Password Link"}

-

{getInvitationUrl()}

-
-
- toast.success("Copied!")}> - - +
+
+

User ID

+

{invitationLinkData?.user_id}

+
+
+ + + + + copyToClipboard(invitationUrl)} + > + + Copy + + + +
diff --git a/ui/litellm-dashboard/src/utils/dataUtils.ts b/ui/litellm-dashboard/src/utils/dataUtils.ts index 8908041a626..c813a69b999 100644 --- a/ui/litellm-dashboard/src/utils/dataUtils.ts +++ b/ui/litellm-dashboard/src/utils/dataUtils.ts @@ -92,6 +92,7 @@ export const copyToClipboard = async ( // Fallback method using document.execCommand (deprecated but widely supported) const fallbackCopyToClipboard = (text: string, messageText: string): boolean => { try { + const previouslyFocused = document.activeElement; const textArea = document.createElement("textarea"); textArea.value = text; @@ -107,6 +108,7 @@ const fallbackCopyToClipboard = (text: string, messageText: string): boolean => const successful = document.execCommand("copy"); document.body.removeChild(textArea); + if (previouslyFocused instanceof HTMLElement) previouslyFocused.focus(); if (successful) { toast.success(messageText); From cb17588276e7f54ec09196e14ef8609dbdc3b240 Mon Sep 17 00:00:00 2001 From: moe-berri Date: Sat, 3 Oct 2026 16:33:45 -0700 Subject: [PATCH 04/18] feat(lens): coordinate worker releases and bundled installs (#44428) * feat(lens): coordinate worker versions and bundled installs * test(lens): exercise bundled Compose startup and restart in CI * fix(lens): refund failed model requests without a response * test(lens): verify trace persistence in the bundled stack * fix(lens): align Helm images and isolate Compose storage * fix(lens): reject worker builds without release identity * fix(lens): encode Compose credentials and normalize worker versions * fix(lens): refuse worker recommendations for unidentified builds --- .github/workflows/image-scan.yml | 10 +- .github/workflows/lens-worker.yml | 9 +- Dockerfile | 2 + backend/Dockerfile | 2 + deploy/lens/Dockerfile | 5 +- deploy/lens/README.md | 72 ++++++++- deploy/lens/compose.build.yaml | 2 + deploy/lens/compose.yaml | 2 +- deploy/lens/config.yaml | 7 + deploy/lens/stack.yaml | 91 +++++++++++ docker/Dockerfile.database | 2 + docker/Dockerfile.non_root | 2 + helm/litellm/templates/_helpers.tpl | 7 + .../litellm/templates/backend/deployment.yaml | 2 + helm/litellm/templates/lens/deployment.yaml | 72 +++++++++ helm/litellm/tests/lens_worker_tests.yaml | 114 ++++++++++++++ helm/litellm/values.yaml | 22 +++ litellm/proxy/_lazy_openapi_snapshot.json | 5 + litellm/proxy/lens/endpoints.py | 22 ++- litellm/proxy/lens/models.py | 1 + litellm/proxy/lens/release.py | 35 +++++ litellm/proxy/lens/worker.py | 33 ++-- scripts/lens_dev.sh | 6 +- tests/e2e/migrations/lens_compose_smoke.sh | 143 ++++++++++++++++++ tests/proxy_behavior/lens/test_lifecycle.py | 8 +- tests/unit/proxy/lens/test_endpoints.py | 47 +++++- tests/unit/proxy/lens/test_release.py | 74 +++++++++ tests/unit/proxy/lens/test_worker.py | 19 +++ tests/unit/test_lens_dev.py | 14 ++ .../worker/WorkerDialog.integration.test.tsx | 7 +- .../lens/setup/worker/WorkerInstall.tsx | 27 +++- .../lens/setup/worker/workerCommand.ts | 9 +- .../lens/setup/worker/workerSchema.test.ts | 11 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 3 + 34 files changed, 841 insertions(+), 46 deletions(-) create mode 100644 deploy/lens/config.yaml create mode 100644 deploy/lens/stack.yaml create mode 100644 helm/litellm/templates/lens/deployment.yaml create mode 100644 helm/litellm/tests/lens_worker_tests.yaml create mode 100644 litellm/proxy/lens/release.py create mode 100644 tests/e2e/migrations/lens_compose_smoke.sh create mode 100644 tests/unit/proxy/lens/test_release.py diff --git a/.github/workflows/image-scan.yml b/.github/workflows/image-scan.yml index 0695720733f..8fb11f40370 100644 --- a/.github/workflows/image-scan.yml +++ b/.github/workflows/image-scan.yml @@ -15,6 +15,9 @@ on: - gateway/main.py - backend/Dockerfile - backend/main.py + - deploy/lens/** + - litellm/proxy/lens/release.py + - tests/e2e/migrations/lens_compose_smoke.sh - docker/component_entrypoint.sh - docker/entrypoint.sh - litellm/proxy/prisma_migration.py @@ -113,7 +116,7 @@ jobs: persist-credentials: false - name: Build runtime image - run: docker build -f Dockerfile -t litellm-runtime-scan:${{ github.sha }} . + run: docker build --build-arg LITELLM_RELEASE_TAG=v0.0.0-lens-ci -f Dockerfile -t litellm-runtime-scan:${{ github.sha }} . - name: Set up Python uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 @@ -127,6 +130,11 @@ jobs: python -m pip install "pytest==9.0.3" python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v + - name: Verify the bundled Lens Compose installation and restart + env: + LITELLM_IMAGE: litellm-runtime-scan:${{ github.sha }} + run: bash tests/e2e/migrations/lens_compose_smoke.sh + migrations-image: name: migrations-image runs-on: ubuntu-latest diff --git a/.github/workflows/lens-worker.yml b/.github/workflows/lens-worker.yml index 54ec2593ed8..e3c47f35c36 100644 --- a/.github/workflows/lens-worker.yml +++ b/.github/workflows/lens-worker.yml @@ -34,7 +34,14 @@ jobs: with: persist-credentials: false - name: Build Lens worker - run: docker build -f deploy/lens/Dockerfile -t lens-worker:${{ github.sha }} . + run: docker build --build-arg LITELLM_RELEASE_TAG=sha-${{ github.sha }} -f deploy/lens/Dockerfile -t lens-worker:${{ github.sha }} . + - name: Reject custom builds without a matching release tag + run: | + if docker build --progress plain -f deploy/lens/Dockerfile -t lens-worker:unversioned . > missing-tag.log 2>&1; then + echo "::error::An unversioned worker build unexpectedly succeeded" + exit 1 + fi + grep -F 'LITELLM_RELEASE_TAG: Pass --build-arg LITELLM_RELEASE_TAG matching the gateway' missing-tag.log - name: Verify standalone imports with a read-only filesystem run: | docker run --rm --network none --read-only --cap-drop ALL --tmpfs /tmp:rw,noexec,nosuid,size=1g \ diff --git a/Dockerfile b/Dockerfile index 4dcecf3ea3d..35dbaa4d41b 100644 --- a/Dockerfile +++ b/Dockerfile @@ -116,6 +116,8 @@ RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \ # Runtime stage FROM $LITELLM_RUNTIME_IMAGE AS runtime +ARG LITELLM_RELEASE_TAG="" +ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG} USER root diff --git a/backend/Dockerfile b/backend/Dockerfile index 59f836b55f8..dfff6e71a46 100644 --- a/backend/Dockerfile +++ b/backend/Dockerfile @@ -71,6 +71,8 @@ RUN sed -i 's/\r$//' docker/component_entrypoint.sh && chmod +x docker/component # ---------- Runtime ---------- FROM $LITELLM_RUNTIME_IMAGE AS runtime +ARG LITELLM_RELEASE_TAG="" +ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG} USER root diff --git a/deploy/lens/Dockerfile b/deploy/lens/Dockerfile index f684940e9a8..3d7ddbe832f 100644 --- a/deploy/lens/Dockerfile +++ b/deploy/lens/Dockerfile @@ -1,7 +1,10 @@ FROM python:3.12-slim +ARG LITELLM_RELEASE_TAG="" +RUN : "${LITELLM_RELEASE_TAG:?Pass --build-arg LITELLM_RELEASE_TAG matching the gateway}" +ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG} WORKDIR /app RUN pip install --no-cache-dir httpx==0.28.1 pydantic==2.11.7 -COPY litellm/proxy/lens/__init__.py litellm/proxy/lens/models.py litellm/proxy/lens/trace_store.py litellm/proxy/lens/analysis.py litellm/proxy/lens/worker.py /app/lens/ +COPY litellm/proxy/lens/__init__.py litellm/proxy/lens/models.py litellm/proxy/lens/trace_store.py litellm/proxy/lens/analysis.py litellm/proxy/lens/worker.py litellm/proxy/lens/release.py /app/lens/ COPY litellm/proxy/lens/prompts/ /app/lens/prompts/ USER 65532:65532 CMD ["python", "-m", "lens.worker"] diff --git a/deploy/lens/README.md b/deploy/lens/README.md index 4f78b7bfb59..afd2f409e21 100644 --- a/deploy/lens/README.md +++ b/deploy/lens/README.md @@ -2,7 +2,60 @@ Lens reviews recorded activity and saves evidence-linked findings in the LiteLLM dashboard under Observability, Lens (`/ui/lens/`) -## Start a worker +## Install the release stack + +Each stable, RC, and dev release containing Lens publishes the worker at the same version on GHCR and Docker Hub. Use the [LiteLLM releases page](https://github.com/BerriAI/litellm/releases) to select a version that includes the coordinated worker release + +For a new local installation, install Docker with Compose, download the two release files, and create a private environment file. Replace `X.Y.Z` with the release version, without `v` (RCs use `X.Y.Z-rc.N`) + +```bash +mkdir litellm-lens +cd litellm-lens +LENS_RELEASE=X.Y.Z +curl -fSLo compose.yaml "https://raw.githubusercontent.com/BerriAI/litellm/v${LENS_RELEASE}/deploy/lens/stack.yaml" +curl -fSLo config.yaml "https://raw.githubusercontent.com/BerriAI/litellm/v${LENS_RELEASE}/deploy/lens/config.yaml" +umask 077 +printf 'LITELLM_VERSION=%s\nLITELLM_MASTER_KEY=sk-%s\nLITELLM_SALT_KEY=sk-%s\n' \ + "$LENS_RELEASE" "$(openssl rand -hex 32)" "$(openssl rand -hex 32)" > .env +printf 'POSTGRES_PASSWORD=%s\nCLICKHOUSE_PASSWORD=%s\n' \ + "$(openssl rand -hex 32)" "$(openssl rand -hex 32)" >> .env +docker compose up -d +``` + +Open `http://localhost:4000/ui/`, log in as `admin` with `LITELLM_MASTER_KEY` from `.env`, and add a model in the dashboard. In Lens, select **Connect worker**, choose that model and a monthly budget, then **Get install command**. Expand **Using Docker Compose or Helm?**, copy the worker token, and add `LENS_WORKER_TOKEN=` to `.env` + +```bash +docker compose --profile lens up -d +``` + +The stack starts LiteLLM, PostgreSQL, ClickHouse, and the worker from published images. The dashboard shows **Worker connected**. The worker has a limited token, no database credentials, and no provider keys. The stack exposes only the dashboard on localhost; use your normal ingress and managed databases for a public production deployment + +Keep `.env` private and preserve its salt key. Keep both named database volumes. To upgrade, wait for active investigations to finish, stop the worker, change only `LITELLM_VERSION`, then pull and recreate the stack: + +```bash +docker compose --profile lens stop lens-worker +# Update LITELLM_VERSION in .env to the new release +docker compose --profile lens pull +docker compose --profile lens up -d +``` + +This preserves your investigations, findings, model credentials, and worker token. Never use `down -v` during an upgrade. If moving from an existing installation, keep its databases and add the standalone worker instead of creating an empty replacement stack + +## Helm + +The componentized `helm/litellm` chart includes an optional Lens worker. Configure PostgreSQL and ClickHouse as usual, install the chart, then obtain a limited worker token from Lens setup. Store it in a Kubernetes Secret and enable the worker in your values: + +```yaml +lensWorker: + enabled: true + tokenSecret: + name: litellm-lens-worker + key: token +``` + +The worker image defaults to the chart's application version, and the chart connects it to the backend service. Keep these values and the Secret when upgrading the chart so the gateway and worker upgrade together. `lensWorker.replicaCount` controls simultaneous investigations. To use a private registry or external proxy, set `lensWorker.image.repository`, `lensWorker.image.tag`, and `lensWorker.url`. The dashboard uses the chart's worker image for standalone install commands too + +## Standalone worker Upgrade your existing LiteLLM proxy to a release that includes Lens with PostgreSQL and agent tracing. Configure one ClickHouse URL for trace writes, bounded reads, and Lens queries: @@ -23,17 +76,17 @@ In **Lens > Investigations**, click **Connect worker**, choose an analysis model The command already contains the compatible worker image and one worker token. The selected virtual key stays on the proxy; its secret is never sent to the worker. No source checkout, environment file, or second LiteLLM deployment is needed. Keep the command private because it includes the token. The LiteLLM release provides the dashboard and APIs; the container only runs background analysis -The dashboard and Compose file pin a verified worker image by digest. The image uses Linux amd64, and the generated command selects that platform. CI also publishes immutable `:sha-` tags for successful worker builds on `main`. Keep the worker image compatible with your gateway version +The dashboard selects the worker image matching the running gateway release. Release images support Linux amd64 and arm64. CI also publishes `:sha-` development images; use those only with a gateway built from the same commit and release tag After upgrading the gateway, update the worker image and redeploy it while keeping its proxy URL and token. Existing containers do not update automatically. If an investigation reports a worker compatibility error, update the image before retrying -For deployments managed with Compose, download `compose.yaml` and provide `LITELLM_URL` and `LENS_WORKER_TOKEN` in an environment file. Its default image is already selected: +For deployments managed with Compose, download `compose.yaml` and provide `LITELLM_URL`, `LENS_WORKER_TOKEN`, and `LITELLM_VERSION` (without `v`) in a private environment file. To use another registry, set `LENS_WORKER_IMAGE` to the compatible image instead of setting a version: ```bash docker compose --env-file /path/to/lens.env -f compose.yaml up -d ``` -Developers can build locally with `LENS_WORKER_IMAGE=litellm-lens-worker:local docker compose -f deploy/lens/compose.yaml -f deploy/lens/compose.build.yaml up -d --build`. To work on Lens itself, `make lens-dev` runs the proxy, a worker from source and the hot-reload dashboard together; set `LENS_DEV_PROXY_PORT` / `LENS_DEV_UI_PORT` to move them off 4000/3000 +To work on Lens itself, `make lens-dev` runs the proxy, a worker from source and the hot-reload dashboard together; set `LENS_DEV_PROXY_PORT` / `LENS_DEV_UI_PORT` to move them off 4000/3000. For a local container build, set `LENS_WORKER_IMAGE=litellm-lens-worker:local` and `LITELLM_RELEASE_TAG` to the gateway's release tag, then use `docker compose -f deploy/lens/compose.yaml -f deploy/lens/compose.build.yaml up -d --build` The generated command gives the worker 1 GiB of temporary memory-backed storage, shared across parallel reviews. Change `size=1g` in the Docker command or set `LENS_WORKER_TMP_SIZE` with Compose to fit your server and workload. A storage failure marks the scan as failed, cleans up temporary traces, and leaves the worker available for other scans; it does not silently truncate the review. Existing workers must be recreated with the new image and mount options @@ -130,3 +183,14 @@ The Lens API now uses `/lens` instead of `/engine`, list responses use `lenses`, Stop workers and let active scans finish before upgrading. Deploy proxy instances together: older proxies cannot use the renamed database tables. The schema migration renames the three Lens tables and the run-history identifier column in place, preserving saved investigations, findings, history, worker credentials, and billing assignments. Existing migration files retain their original names and checksums Upgrades using `--use_prisma_db_push` stop before schema changes if any legacy Lens table exists, preventing Prisma from dropping saved data. Apply `litellm-proxy-extras/litellm_proxy_extras/migrations/20261001100000_rename_lens/migration.sql` to the configured database schema before retrying. Deployments already using migration history can instead start without `--use_prisma_db_push` to apply the shipped migration normally. Fresh databases and databases already using the renamed tables can continue using database push + + +## Release compatibility + +Released gateway and worker images carry `LITELLM_RELEASE_TAG`. A worker announces its release and protocol before claiming an investigation. A mismatch returns HTTP 409 with the required image, leaving queued investigations untouched. During a rolling upgrade, workers wait for a gateway from their release + +The dashboard reads its image from the running gateway. `LENS_WORKER_IMAGE` overrides the registry/image for private deployments. Worker-only Compose accepts `LITELLM_VERSION` (without `v`) or an explicit `LENS_WORKER_IMAGE`. Release workers are available as `ghcr.io/berriai/litellm-lens-worker:vX.Y.Z` and `docker.io/litellm/litellm-lens-worker:vX.Y.Z`, including matching RC/dev suffixes, on amd64 and arm64 + +For source development, use `make lens-dev`, which gives the proxy and source worker the same commit identity. For custom containers, build both from the same checkout with `--build-arg LITELLM_RELEASE_TAG=sha-$(git rev-parse HEAD)` and set the proxy's `LENS_WORKER_IMAGE` to the worker image you built. An unlabelled custom build refuses worker setup and claims instead of guessing from the Python package version. Normal package-index installations use their installed release version + +The hourly development pipeline pins all component images to the same selected commit and publishes its chart only after every build and worker smoke test succeeds. The public commit-tagged worker workflow publishes on Lens-related changes, so an arbitrary `main` commit may require building your own pair; do not substitute the newest available worker diff --git a/deploy/lens/compose.build.yaml b/deploy/lens/compose.build.yaml index e4237d8de23..52d59a84a79 100644 --- a/deploy/lens/compose.build.yaml +++ b/deploy/lens/compose.build.yaml @@ -3,4 +3,6 @@ services: build: context: ../.. dockerfile: deploy/lens/Dockerfile + args: + LITELLM_RELEASE_TAG: ${LITELLM_RELEASE_TAG:?Set the release tag used by the gateway} image: litellm-lens-worker:local diff --git a/deploy/lens/compose.yaml b/deploy/lens/compose.yaml index fc9850fcb04..799f0a4fb1e 100644 --- a/deploy/lens/compose.yaml +++ b/deploy/lens/compose.yaml @@ -1,6 +1,6 @@ services: lens-worker: - image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker@sha256:44f0597c7583dcfef999ece9a8bc02cfeb9f0f5167a1221cee3bd10b1b79271b} + image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker:v${LITELLM_VERSION:?Set LITELLM_VERSION to the gateway release, without the v prefix}} environment: LITELLM_URL: ${LITELLM_URL:?Set the URL reachable from this container} LENS_WORKER_TOKEN: ${LENS_WORKER_TOKEN:?Create a worker credential in the Lens UI} diff --git a/deploy/lens/config.yaml b/deploy/lens/config.yaml new file mode 100644 index 00000000000..cb12a2b0919 --- /dev/null +++ b/deploy/lens/config.yaml @@ -0,0 +1,7 @@ +general_settings: + master_key: os.environ/LITELLM_MASTER_KEY + tracing: + store: + type: clickhouse + url: os.environ/CLICKHOUSE_URL + retention_days: 14 diff --git a/deploy/lens/stack.yaml b/deploy/lens/stack.yaml new file mode 100644 index 00000000000..ab559e27b19 --- /dev/null +++ b/deploy/lens/stack.yaml @@ -0,0 +1,91 @@ +name: litellm-lens + +services: + litellm: + image: ghcr.io/berriai/litellm:${LITELLM_VERSION:?Set LITELLM_VERSION to a published release, without the v prefix} + entrypoint: + - python3 + - -c + - | + import os, sys + from urllib.parse import quote + postgres_password = quote(os.environ["POSTGRES_PASSWORD"], safe="") + clickhouse_password = quote(os.environ["CLICKHOUSE_PASSWORD"], safe="") + os.environ["DATABASE_URL"] = f"postgresql://litellm:{postgres_password}@db:5432/litellm" + os.environ["CLICKHOUSE_URL"] = f"http://default:{clickhouse_password}@clickhouse:8123" + os.execv("docker/prod_entrypoint.sh", ["docker/prod_entrypoint.sh", *sys.argv[1:]]) + command: ["--config", "/app/lens-config.yaml", "--port", "4000"] + environment: + LITELLM_MASTER_KEY: ${LITELLM_MASTER_KEY:?Set a strong master key} + LITELLM_SALT_KEY: ${LITELLM_SALT_KEY:?Set a permanent encryption key and keep it across upgrades} + POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:?Set a permanent database password} + STORE_MODEL_IN_DB: "True" + CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD:?Set a permanent ClickHouse password} + LENS_WORKER_IMAGE: ghcr.io/berriai/litellm-lens-worker:v${LITELLM_VERSION} + volumes: + - ./config.yaml:/app/lens-config.yaml:ro + ports: + - "127.0.0.1:${LITELLM_PORT:-4000}:4000" + networks: [proxy, storage] + depends_on: + db: + condition: service_healthy + clickhouse: + condition: service_healthy + restart: unless-stopped + + lens-worker: + profiles: [lens] + image: ghcr.io/berriai/litellm-lens-worker:v${LITELLM_VERSION} + environment: + LITELLM_URL: http://litellm:4000 + LENS_WORKER_TOKEN: ${LENS_WORKER_TOKEN:-} + depends_on: [litellm] + networks: [proxy] + restart: unless-stopped + read_only: true + tmpfs: + - /tmp:rw,noexec,nosuid,size=${LENS_WORKER_TMP_SIZE:-1g} + cap_drop: [ALL] + security_opt: [no-new-privileges:true] + + db: + image: postgres:16 + environment: + POSTGRES_DB: litellm + POSTGRES_USER: litellm + POSTGRES_PASSWORD: ${POSTGRES_PASSWORD} + networks: [storage] + volumes: + - postgres_data:/var/lib/postgresql/data + healthcheck: + test: ["CMD-SHELL", "pg_isready -U litellm -d litellm"] + interval: 5s + timeout: 5s + retries: 20 + restart: unless-stopped + + clickhouse: + image: clickhouse/clickhouse-server:26.9.6.6 + environment: + CLICKHOUSE_USER: default + CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD} + CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT: "1" + volumes: + - clickhouse_data:/var/lib/clickhouse + healthcheck: + test: ["CMD", "clickhouse-client", "--user", "default", "--password", "${CLICKHOUSE_PASSWORD}", "--query", "SELECT 1"] + interval: 5s + timeout: 5s + retries: 20 + restart: unless-stopped + networks: [storage] + +networks: + proxy: + storage: + internal: true + +volumes: + postgres_data: + clickhouse_data: diff --git a/docker/Dockerfile.database b/docker/Dockerfile.database index 61b6faae691..3309fdd5341 100644 --- a/docker/Dockerfile.database +++ b/docker/Dockerfile.database @@ -113,6 +113,8 @@ RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \ sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh FROM $LITELLM_RUNTIME_IMAGE AS runtime +ARG LITELLM_RELEASE_TAG="" +ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG} USER root diff --git a/docker/Dockerfile.non_root b/docker/Dockerfile.non_root index ca526e06834..bafd1af46d1 100644 --- a/docker/Dockerfile.non_root +++ b/docker/Dockerfile.non_root @@ -122,6 +122,8 @@ RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \ sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh FROM $LITELLM_RUNTIME_IMAGE AS runtime +ARG LITELLM_RELEASE_TAG="" +ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG} WORKDIR /app USER root diff --git a/helm/litellm/templates/_helpers.tpl b/helm/litellm/templates/_helpers.tpl index 20fd1a722dc..2b7fdfb5fd7 100644 --- a/helm/litellm/templates/_helpers.tpl +++ b/helm/litellm/templates/_helpers.tpl @@ -471,6 +471,13 @@ Directory of the collector's unix socket, shared by the gateway and collector containers through an emptyDir. Empty when the sidecar is off or gateway.collector.address is a tcp://127.0.0.1: address. */}} +{{- define "litellm.lensWorker.image" -}} +{{- $backendTag := .Values.backend.image.tag | default .Chart.AppVersion -}} +{{- $releaseTag := ternary (printf "v%s" $backendTag) $backendTag (regexMatch "^[0-9]" $backendTag) -}} +{{- $tag := .Values.lensWorker.image.tag | default $releaseTag -}} +{{- printf "%s:%s" .Values.lensWorker.image.repository $tag -}} +{{- end -}} + {{- define "litellm.gateway.collectorSocketDir" -}} {{- if and .Values.gateway.collector.enabled (hasPrefix "unix://" .Values.gateway.collector.address) -}} {{- dir (trimPrefix "unix://" .Values.gateway.collector.address) -}} diff --git a/helm/litellm/templates/backend/deployment.yaml b/helm/litellm/templates/backend/deployment.yaml index 3eb64e5528c..5d3be1439bd 100644 --- a/helm/litellm/templates/backend/deployment.yaml +++ b/helm/litellm/templates/backend/deployment.yaml @@ -57,6 +57,8 @@ spec: containerPort: 4001 protocol: TCP env: + - name: LENS_WORKER_IMAGE + value: {{ include "litellm.lensWorker.image" . | quote }} {{- include "litellm.serverEnv" (dict "root" $ "component" .Values.backend) | nindent 12 }} {{- if .Values.gateway.config.create }} - name: CONFIG_FILE_PATH diff --git a/helm/litellm/templates/lens/deployment.yaml b/helm/litellm/templates/lens/deployment.yaml new file mode 100644 index 00000000000..787581b9ad1 --- /dev/null +++ b/helm/litellm/templates/lens/deployment.yaml @@ -0,0 +1,72 @@ +{{- if .Values.lensWorker.enabled }} +apiVersion: apps/v1 +kind: Deployment +metadata: + name: {{ include "litellm.fullname" . }}-lens-worker + labels: + {{- include "litellm.commonLabels" . | nindent 4 }} + app.kubernetes.io/component: lens-worker +spec: + replicas: {{ .Values.lensWorker.replicaCount }} + selector: + matchLabels: + app.kubernetes.io/instance: {{ .Release.Name }} + app.kubernetes.io/component: lens-worker + template: + metadata: + labels: + {{- include "litellm.commonLabels" . | nindent 8 }} + app.kubernetes.io/component: lens-worker + spec: + automountServiceAccountToken: false + {{- with .Values.imagePullSecrets }} + imagePullSecrets: + {{- toYaml . | nindent 8 }} + {{- end }} + securityContext: + runAsNonRoot: true + runAsUser: 65532 + runAsGroup: 65532 + fsGroup: 65532 + seccompProfile: + type: RuntimeDefault + containers: + - name: lens-worker + image: {{ include "litellm.lensWorker.image" . | quote }} + imagePullPolicy: {{ .Values.lensWorker.image.pullPolicy }} + securityContext: + allowPrivilegeEscalation: false + readOnlyRootFilesystem: true + capabilities: + drop: [ALL] + env: + - name: LITELLM_URL + value: {{ .Values.lensWorker.url | default (printf "http://%s:%v" (include "litellm.backend.fullname" .) .Values.backend.service.port) | quote }} + - name: LENS_WORKER_TOKEN + valueFrom: + secretKeyRef: + name: {{ required "lensWorker.tokenSecret.name must reference a Lens worker token" .Values.lensWorker.tokenSecret.name | quote }} + key: {{ .Values.lensWorker.tokenSecret.key | quote }} + resources: + {{- toYaml .Values.lensWorker.resources | nindent 12 }} + volumeMounts: + - name: tmp + mountPath: /tmp + volumes: + - name: tmp + emptyDir: + medium: Memory + sizeLimit: {{ .Values.lensWorker.tmpSizeLimit }} + {{- with .Values.lensWorker.nodeSelector }} + nodeSelector: + {{- toYaml . | nindent 8 }} + {{- end }} + {{- with .Values.lensWorker.tolerations }} + tolerations: + {{- toYaml . | nindent 8 }} + {{- end }} + {{- with .Values.lensWorker.affinity }} + affinity: + {{- toYaml . | nindent 8 }} + {{- end }} +{{- end }} diff --git a/helm/litellm/tests/lens_worker_tests.yaml b/helm/litellm/tests/lens_worker_tests.yaml new file mode 100644 index 00000000000..5230f14efb4 --- /dev/null +++ b/helm/litellm/tests/lens_worker_tests.yaml @@ -0,0 +1,114 @@ +suite: Lens worker release and credentials +templates: + - lens/deployment.yaml + - backend/deployment.yaml + - gateway/configmap.yaml +values: + - ./values/required.yaml +tests: + - it: keeps the worker opt in + template: lens/deployment.yaml + asserts: + - hasDocuments: + count: 0 + - it: requires a limited worker credential when enabled + template: lens/deployment.yaml + set: + lensWorker.enabled: true + asserts: + - failedTemplate: + errorMessage: lensWorker.tokenSecret.name must reference a Lens worker token + - it: uses the chart release and a secret without granting Kubernetes access + template: lens/deployment.yaml + chart: + appVersion: v1.2.3 + set: + lensWorker.enabled: true + lensWorker.tokenSecret.name: lens-credential + asserts: + - equal: + path: spec.template.spec.containers[0].image + value: ghcr.io/berriai/litellm-lens-worker:v1.2.3 + - equal: + path: spec.template.spec.containers[0].env[1].valueFrom.secretKeyRef + value: + name: lens-credential + key: token + - equal: + path: spec.template.spec.automountServiceAccountToken + value: false + - equal: + path: spec.template.spec.containers[0].securityContext.readOnlyRootFilesystem + value: true + - equal: + path: spec.template.spec.volumes[0].emptyDir + value: + medium: Memory + sizeLimit: 1Gi + - it: advertises the same private dev image to standalone installers + template: backend/deployment.yaml + set: + lensWorker.image.repository: registry.example/lens-worker + lensWorker.image.tag: branch-main-1234567 + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LENS_WORKER_IMAGE + value: registry.example/lens-worker:branch-main-1234567 + - it: supports an external gateway and a registry override + template: lens/deployment.yaml + set: + lensWorker.enabled: true + lensWorker.tokenSecret.name: lens-credential + lensWorker.url: https://gateway.example/proxy + lensWorker.image.repository: registry.example/lens-worker + lensWorker.image.tag: branch-main-1234567 + asserts: + - equal: + path: spec.template.spec.containers[0].image + value: registry.example/lens-worker:branch-main-1234567 + - equal: + path: spec.template.spec.containers[0].env[0].value + value: https://gateway.example/proxy + - it: prefixes a numeric chart release with v + template: lens/deployment.yaml + chart: + appVersion: 1.2.3-rc.4 + set: + lensWorker.enabled: true + lensWorker.tokenSecret.name: lens-credential + asserts: + - equal: + path: spec.template.spec.containers[0].image + value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-rc.4 + - it: follows a backend image override when no worker tag is set + template: lens/deployment.yaml + set: + backend.image.tag: branch-main-1234567 + lensWorker.enabled: true + lensWorker.tokenSecret.name: lens-credential + asserts: + - equal: + path: spec.template.spec.containers[0].image + value: ghcr.io/berriai/litellm-lens-worker:branch-main-1234567 + - it: recommends the overridden backend release for standalone installers + template: backend/deployment.yaml + set: + backend.image.tag: v1.2.3-dev.4 + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: LENS_WORKER_IMAGE + value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-dev.4 + - it: normalizes a numeric backend tag to the published worker tag + template: lens/deployment.yaml + set: + backend.image.tag: 1.2.3-dev.4 + lensWorker.enabled: true + lensWorker.tokenSecret.name: lens-credential + asserts: + - equal: + path: spec.template.spec.containers[0].image + value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-dev.4 diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index 2c0c7151a32..6964dfd2b6e 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -629,3 +629,25 @@ ui: affinity: {} # Same shape as gateway.topologySpreadConstraints. topologySpreadConstraints: [] + +lensWorker: + enabled: false + replicaCount: 1 + image: + repository: ghcr.io/berriai/litellm-lens-worker + tag: "" + pullPolicy: IfNotPresent + tokenSecret: + name: "" + key: token + url: "" + tmpSizeLimit: 1Gi + resources: + requests: + cpu: 100m + memory: 256Mi + limits: + memory: 2Gi + nodeSelector: {} + tolerations: [] + affinity: {} diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 7b9af2f2b98..a927d4b415f 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -34885,6 +34885,10 @@ "WorkerCreated": { "additionalProperties": false, "properties": { + "image": { + "title": "Image", + "type": "string" + }, "token": { "title": "Token", "type": "string" @@ -34894,6 +34898,7 @@ } }, "required": [ + "image", "worker", "token" ], diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index b6c033fd17b..3ae9db0a600 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -43,6 +43,7 @@ from litellm.proxy.lens.models import ( Worker, WorkerCreated, ) +from litellm.proxy.lens.release import PROTOCOL_VERSION, release_tag, worker_image from litellm.proxy.lens.repository import LensRepository, WriterDatabase from litellm.proxy.lens.sources import ActivityAvailability, SourceReader, Storage, parse_execution from litellm.proxy.lens.state import ( @@ -421,9 +422,20 @@ class WorkerName(WorkerBilling): name: str = Field(default="Lens worker", min_length=1) +def configured_worker_image() -> str: + if image := worker_image(): + return image + raise HTTPException( + 503, + "This LiteLLM build has no release identity. Use a published release, make lens-dev, " + "or build the gateway and worker from the same commit with the same LITELLM_RELEASE_TAG.", + ) + + @router.post("/workers/register", response_model=WorkerCreated) async def register_worker(body: WorkerName, auth: Auth) -> WorkerCreated: scope: Final = user_scope(auth, write=True) + image: Final = configured_worker_image() await validate_key(body.analysis_key_id) token: Final = "lens-" + secrets.token_urlsafe(40) worker: Final = Worker( @@ -434,7 +446,7 @@ async def register_worker(body: WorkerName, auth: Auth) -> WorkerCreated: last_seen=datetime(1970, 1, 1, tzinfo=timezone.utc), ) await repository().save_worker(worker, hashlib.sha256(token.encode()).hexdigest()) - return WorkerCreated(worker=worker, token=token) + return WorkerCreated(worker=worker, token=token, image=image) @router.put("/workers/{worker_id}/billing-key", response_model=Worker) @@ -466,9 +478,11 @@ async def revoke_worker(worker_id: str, auth: Auth) -> bool: @router.post("/worker/claim", response_model=Claim | None) -async def claim(worker: WorkerAuth, protocol_version: int = 1) -> Claim | None: - if protocol_version not in (2, 3): - raise HTTPException(409, "Upgrade the Lens worker using the current Connect worker command") +async def claim(worker: WorkerAuth, protocol_version: int = 1, worker_release: str = "") -> Claim | None: + image: Final = configured_worker_image() + expected: Final = release_tag() + if protocol_version != PROTOCOL_VERSION or worker_release != expected: + raise HTTPException(409, f"Upgrade the Lens worker to {image} and retry") if worker.analysis_key_id is None: raise HTTPException(409, "Assign an analysis key to this worker in Lens setup") now: Final = datetime.now(timezone.utc) diff --git a/litellm/proxy/lens/models.py b/litellm/proxy/lens/models.py index de39dd9de74..e3a08103c8d 100644 --- a/litellm/proxy/lens/models.py +++ b/litellm/proxy/lens/models.py @@ -252,6 +252,7 @@ class Worker(Record): class WorkerCreated(Record): + image: str worker: Worker token: str diff --git a/litellm/proxy/lens/release.py b/litellm/proxy/lens/release.py new file mode 100644 index 00000000000..0338858a6a5 --- /dev/null +++ b/litellm/proxy/lens/release.py @@ -0,0 +1,35 @@ +import os +from importlib.metadata import PackageNotFoundError, distribution +from pathlib import Path +from typing import Final + +PROTOCOL_VERSION: Final = 4 + + +def release_tag() -> str: + if "LITELLM_RELEASE_TAG" in os.environ: + return os.environ["LITELLM_RELEASE_TAG"] + try: + installed: Final = distribution("litellm") + except PackageNotFoundError: + return "" + if installed.read_text("direct_url.json") is not None: + return "" + if Path(str(installed.locate_file("litellm/proxy/lens/release.py"))).resolve() != Path(__file__).resolve(): + return "" + + from packaging.version import Version + + parsed: Final = Version(installed.version) + suffix: Final = f"-dev.{parsed.dev}" if parsed.dev is not None else f"-rc.{parsed.pre[1]}" if parsed.pre else "" + return f"v{parsed.base_version}{suffix}" + + +def worker_image() -> str: + tag: Final = release_tag() + if not tag: + return "" + override: Final = os.environ.get("LENS_WORKER_IMAGE", "") + if override: + return override + return f"ghcr.io/berriai/litellm-lens-worker:{tag}" diff --git a/litellm/proxy/lens/worker.py b/litellm/proxy/lens/worker.py index f9e746489d9..06455dc1a9a 100644 --- a/litellm/proxy/lens/worker.py +++ b/litellm/proxy/lens/worker.py @@ -11,6 +11,7 @@ from pydantic import BaseModel, ConfigDict, ValidationError from .analysis import AnalysisResponseError, analyze_sample, validation_details from .models import Claim, Coverage, ExecutionContent, ModelRequest, ModelResult, Progress, Result, Sample +from .release import PROTOCOL_VERSION, release_tag logger: Final = logging.getLogger("litellm.lens.worker") @@ -120,8 +121,26 @@ class LensWorker: await self.sleep(2**attempt) return await self.model_request(path, body, attempt + 1) + async def report_unreadable_claim(self, identity: ClaimIdentity) -> None: + failure: Final = await self.client.post( + f"/lens/worker/{identity.lens_id}/{identity.job.id}/result", + json=Result( + coverage=Coverage(), + error="The worker could not read this investigation. Update the worker to match the gateway, then retry.", + ).model_dump(), + ) + if failure.status_code != 409: + failure.raise_for_status() + logger.warning("Worker could not read a claimed investigation; reported a version compatibility failure") + async def run_once(self) -> bool: - response: Final = await self.client.post("/lens/worker/claim", params=MappingProxyType({"protocol_version": 3})) + response: Final = await self.client.post( + "/lens/worker/claim", + params=MappingProxyType({"protocol_version": str(PROTOCOL_VERSION), "worker_release": release_tag()}), + ) + if response.status_code == 409: + logger.warning("Lens worker cannot claim work: %s", response.text) + return False response.raise_for_status() payload: Final = response.json() if payload is None: @@ -129,17 +148,7 @@ class LensWorker: try: claim: Final = Claim.model_validate(payload) except ValidationError: - identity: Final = ClaimIdentity.model_validate(payload) - failure: Final = await self.client.post( - f"/lens/worker/{identity.lens_id}/{identity.job.id}/result", - json=Result( - coverage=Coverage(), - error="The worker could not read this investigation. Update the worker to match the gateway, then retry.", - ).model_dump(), - ) - if failure.status_code != 409: - failure.raise_for_status() - logger.warning("Worker could not read a claimed investigation; reported a version compatibility failure") + await self.report_unreadable_claim(ClaimIdentity.model_validate(payload)) return True prefix: Final = f"/lens/worker/{claim.lens_id}/{claim.job.id}" diff --git a/scripts/lens_dev.sh b/scripts/lens_dev.sh index 7ab4b9bb433..46408fdaf19 100755 --- a/scripts/lens_dev.sh +++ b/scripts/lens_dev.sh @@ -13,6 +13,7 @@ set -euo pipefail repo_root="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +source_release_tag="sha-$(git -C "$repo_root" rev-parse HEAD)" proxy_port="${LENS_DEV_PROXY_PORT:-4000}" ui_port="${LENS_DEV_UI_PORT:-3000}" state_dir="${LENS_DEV_STATE_DIR:-$repo_root/.lens-dev}" @@ -108,6 +109,8 @@ proxy_env() { unset ANTHROPIC_BASE_URL ANTHROPIC_AUTH_TOKEN ANTHROPIC_CUSTOM_HEADERS OPENAI_BASE_URL OPENAI_API_BASE for var in $(compgen -e | grep '^REDIS_' || true); do unset "$var"; done eval "$1" + export LITELLM_RELEASE_TAG="$source_release_tag" + export LENS_WORKER_IMAGE=litellm-lens-worker:local export LITELLM_MODE=PRODUCTION export LITELLM_MASTER_KEY="$master_key" if [ "$master_key" = sk-1234 ]; then export LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY=true; fi @@ -236,7 +239,8 @@ main() { wait_for_proxy "$proxy_pid" ensure_worker_token - LITELLM_MODE=PRODUCTION LITELLM_URL="$proxy_url" LENS_WORKER_TOKEN="$(cat "$token_file")" \ + LITELLM_RELEASE_TAG="$source_release_tag" \ + LITELLM_MODE=PRODUCTION LITELLM_URL="$proxy_url" LENS_WORKER_TOKEN="$(cat "$token_file")" \ "$py" -c "import asyncio, logging; from litellm.proxy.lens.worker import main; logging.basicConfig(level=logging.INFO); asyncio.run(main())" \ < /dev/null > "$log_dir/worker.log" 2>&1 & pids+=("$!") diff --git a/tests/e2e/migrations/lens_compose_smoke.sh b/tests/e2e/migrations/lens_compose_smoke.sh new file mode 100644 index 00000000000..69a13b58d88 --- /dev/null +++ b/tests/e2e/migrations/lens_compose_smoke.sh @@ -0,0 +1,143 @@ +#!/usr/bin/env bash +set -euo pipefail + +qa_dir=$(mktemp -d) +master_key="sk-$(openssl rand -hex 32)" +compose=(docker compose -p lens-compose-ci --env-file "$qa_dir/env" -f deploy/lens/stack.yaml) +cleanup() { + "${compose[@]}" --profile lens down -v --remove-orphans >/dev/null 2>&1 || true + rm -rf "$qa_dir" +} +trap cleanup EXIT +umask 077 +printf 'LITELLM_VERSION=0.0.0-lens-ci\nLITELLM_PORT=4418\nLITELLM_MASTER_KEY=%s\nLITELLM_SALT_KEY=sk-%s\n' \ + "$master_key" "$(openssl rand -hex 32)" > "$qa_dir/env" +printf 'POSTGRES_PASSWORD=%s:/?#@%%\nCLICKHOUSE_PASSWORD=%s:/?#@%%\n' \ + "$(openssl rand -hex 32)" "$(openssl rand -hex 32)" >> "$qa_dir/env" +docker tag "${LITELLM_IMAGE:?Set LITELLM_IMAGE to the built gateway image}" ghcr.io/berriai/litellm:0.0.0-lens-ci +docker build --build-arg LITELLM_RELEASE_TAG=v0.0.0-lens-ci -f deploy/lens/Dockerfile \ + -t ghcr.io/berriai/litellm-lens-worker:v0.0.0-lens-ci . +"${compose[@]}" up -d + +api() { + curl --fail-with-body --silent --show-error --max-time 30 \ + -H "Authorization: Bearer $master_key" -H 'Content-Type: application/json' \ + "http://127.0.0.1:4418$1" "${@:2}" +} +ready=false +for attempt in $(seq 1 90); do + if api /health/liveliness > /dev/null 2>&1; then ready=true; break; fi + sleep 2 +done +if [[ "$ready" != true ]]; then "${compose[@]}" logs litellm; exit 1; fi +trace_id=$(openssl rand -hex 16) +span_id=$(openssl rand -hex 8) +start_ns="$(date +%s)000000000" +jq -n --arg trace "$trace_id" --arg span "$span_id" --arg at "$start_ns" \ + '{resourceSpans:[{resource:{attributes:[{key:"service.name",value:{stringValue:"lens-compose-ci"}}]}, + scopeSpans:[{scope:{name:"lens-compose-ci"},spans:[{traceId:$trace,spanId:$span,name:"Compose trace", + kind:1,startTimeUnixNano:$at,endTimeUnixNano:$at, + attributes:[{key:"openinference.span.kind",value:{stringValue:"AGENT"}}],status:{code:1}}]}]}]}' \ + > "$qa_dir/trace.json" +api /v1/traces -d "@$qa_dir/trace.json" > /dev/null +trace_saved() { + for attempt in $(seq 1 60); do + if api "/v1/traces/$trace_id" > "$qa_dir/saved-trace.json" 2>/dev/null && \ + jq -e --arg trace "$trace_id" --arg span "$span_id" \ + '.summary.trace_id == $trace and any(.spans[]; .span_id == $span)' "$qa_dir/saved-trace.json" > /dev/null; then + return 0 + fi + sleep 2 + done + return 1 +} +trace_saved +api /key/generate -d '{"key_alias":"Lens Compose CI","models":["lens-compose-ci"],"max_budget":1}' > "$qa_dir/key.json" +key_id=$(jq -r '.token_id // empty' "$qa_dir/key.json") +if [[ -z "$key_id" ]]; then + key_id=$(jq -rj '.key' "$qa_dir/key.json" | openssl dgst -sha256 | awk '{print $NF}') +fi +jq -n --arg key "$key_id" '{name:"Lens Compose CI",analysis_key_id:$key}' > "$qa_dir/registration.json" +api /lens/workers/register -d "@$qa_dir/registration.json" > "$qa_dir/worker.json" +jq -e '.image == "ghcr.io/berriai/litellm-lens-worker:v0.0.0-lens-ci"' "$qa_dir/worker.json" > /dev/null +printf 'LENS_WORKER_TOKEN=%s\n' "$(jq -r '.token' "$qa_dir/worker.json")" >> "$qa_dir/env" +worker_id=$(jq -r '.worker.id' "$qa_dir/worker.json") +heartbeat_after=$(date -u +'%Y-%m-%dT%H:%M:%S') +"${compose[@]}" --profile lens up -d + +connected() { + for attempt in $(seq 1 60); do + if api /lens > "$qa_dir/lens.json" 2>/dev/null && \ + jq -e --arg id "$worker_id" --arg since "$heartbeat_after" \ + '.workers[] | select(.id == $id and .last_seen > $since)' "$qa_dir/lens.json" > /dev/null; then + return 0 + fi + sleep 2 + done + "${compose[@]}" --profile lens logs lens-worker + return 1 +} +connected +printf 'Fresh Compose stack: matching worker image and authenticated heartbeat passed\n' + +for target in db:5432 clickhouse:8123; do + service=${target%:*} + port=${target#*:} + address=$(docker inspect --format '{{range .NetworkSettings.Networks}}{{.IPAddress}}{{end}}' "$("${compose[@]}" ps -q "$service")") + "${compose[@]}" exec -T lens-worker python -c ' +import socket, sys +for host in (sys.argv[1], sys.argv[2]): + try: + connection = socket.create_connection((host, int(sys.argv[3])), timeout=2) + except OSError: + continue + connection.close() + raise SystemExit("Worker can reach a datastore directly") +' "$service" "$address" "$port" +done +printf 'Worker can reach the proxy but cannot connect directly to PostgreSQL or ClickHouse\n' + +status=$(curl --silent --show-error -o "$qa_dir/mismatch.json" -w '%{http_code}' -X POST \ + -H "Authorization: Bearer $(jq -r '.token' "$qa_dir/worker.json")" \ + 'http://127.0.0.1:4418/lens/worker/claim?protocol_version=4&worker_release=v0.0.0-old') +[[ "$status" == 409 ]] +jq -e '.detail | contains("Upgrade the Lens worker")' "$qa_dir/mismatch.json" > /dev/null + +"${compose[@]}" --profile lens restart litellm lens-worker +heartbeat_after=$(date -u +'%Y-%m-%dT%H:%M:%S') +connected +trace_saved +api /lens > "$qa_dir/restarted.json" +jq -e --arg id "$worker_id" --arg key "$key_id" \ + '.workers[] | select(.id == $id and .analysis_key_id == $key)' "$qa_dir/restarted.json" > /dev/null +printf 'Compose restart: trace, worker identity, token and billing assignment preserved; wrong release rejected\n' + +cat > "$qa_dir/unversioned.yaml" <<'EOF' +services: + litellm: + environment: + LITELLM_RELEASE_TAG: "" +EOF +"${compose[@]}" -f "$qa_dir/unversioned.yaml" up -d litellm +for attempt in $(seq 1 90); do + if api /health/liveliness > /dev/null 2>&1; then break; fi + sleep 2 +done +api /health/liveliness > /dev/null +status=$(curl --silent --show-error --max-time 30 -o "$qa_dir/unversioned-registration.json" -w '%{http_code}' \ + -H "Authorization: Bearer $master_key" -H 'Content-Type: application/json' \ + -d "@$qa_dir/registration.json" 'http://127.0.0.1:4418/lens/workers/register') +[[ "$status" == 503 ]] +jq -e '.detail | contains("no release identity")' "$qa_dir/unversioned-registration.json" > /dev/null +status=$(curl --silent --show-error --max-time 30 -o "$qa_dir/unversioned-claim.json" -w '%{http_code}' -X POST \ + -H "Authorization: Bearer $(jq -r '.token' "$qa_dir/worker.json")" \ + 'http://127.0.0.1:4418/lens/worker/claim?protocol_version=4&worker_release=') +[[ "$status" == 503 ]] +jq -e '.detail | contains("no release identity")' "$qa_dir/unversioned-claim.json" > /dev/null +api /lens > "$qa_dir/unversioned-workers.json" +jq -e --arg id "$worker_id" '.workers | length == 1 and .[0].id == $id' "$qa_dir/unversioned-workers.json" > /dev/null +"${compose[@]}" up -d litellm +heartbeat_after=$(date -u +'%Y-%m-%dT%H:%M:%S') +connected +trace_saved +printf 'Unversioned gateway: setup and claims refused without guessing; original worker and trace recovered\n' diff --git a/tests/proxy_behavior/lens/test_lifecycle.py b/tests/proxy_behavior/lens/test_lifecycle.py index 1dade952dd8..e3027474baa 100644 --- a/tests/proxy_behavior/lens/test_lifecycle.py +++ b/tests/proxy_behavior/lens/test_lifecycle.py @@ -29,13 +29,15 @@ from litellm.proxy.lens.models import ( Scope, Worker, ) +from litellm.proxy.lens.release import PROTOCOL_VERSION, release_tag from litellm.proxy.lens.repository import Database, LensRepository, Row from litellm.proxy.lens.state import can_access from litellm.proxy.utils import PrismaClient, ProxyLogging @pytest_asyncio.fixture(loop_scope="function") -async def lens_database() -> AsyncIterator[PrismaClient]: +async def lens_database(monkeypatch: pytest.MonkeyPatch) -> AsyncIterator[PrismaClient]: + monkeypatch.setenv("LITELLM_RELEASE_TAG", "v0.0.0-lens-lifecycle") original_db: Final = proxy_server.prisma_client original_router: Final = proxy_server.llm_router original_settings: Final = proxy_server.general_settings @@ -349,8 +351,9 @@ async def test_scan_lifecycle_persists_results_and_revokes_worker(lens_database: authenticated_legacy: Final = await endpoints.worker_auth(credentials) assert authenticated_legacy.analysis_key_id is None with pytest.raises(HTTPException) as needs_billing: - await endpoints.claim(authenticated_legacy, protocol_version=2) + await endpoints.claim(authenticated_legacy, protocol_version=PROTOCOL_VERSION, worker_release=release_tag()) assert needs_billing.value.status_code == 409 + assert "Assign an analysis key" in needs_billing.value.detail assert await endpoints.heartbeat(lens.id, claimed.job.id, authenticated_legacy) finished: Final = await endpoints.result( lens.id, claimed.job.id, Result(coverage=Coverage(screened=2)), authenticated_legacy, storage=None @@ -441,6 +444,7 @@ async def test_failed_model_requests_release_lens_budget_reservations(lens_datab stored: Final = await endpoints.get_lens(lens.id, worker.scope) assert stored.spent == 0 assert stored.jobs[0].cost == 0 + assert not any(step.kind == "model" for step in stored.jobs[0].steps) finally: await lens_database.db.execute_raw('DELETE FROM "LiteLLM_LensRun" WHERE lens_id=$1', lens.id) await lens_database.db.execute_raw('DELETE FROM "LiteLLM_Lens" WHERE id=$1', lens.id) diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index 783730b7c8f..a0d01878277 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -6,16 +6,16 @@ from fastapi import HTTPException from pydantic import ValidationError import litellm -from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm import Router +from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.lens.endpoints import ( list_agents, run_settings, run_window, user_scope, + validate_model, watchable, watching, - validate_model, worker_supports_model, ) from litellm.proxy.lens.models import Lens, LensSettings, RunRequest, Scope @@ -159,13 +159,17 @@ def test_invalid_explicit_execution_ids_are_rejected(identity: str) -> None: assert error.value.status_code == 422 +@pytest.mark.parametrize("protocol_version", (1, 2, 3)) @pytest.mark.asyncio -async def test_incompatible_worker_is_rejected_before_claiming_work() -> None: +async def test_incompatible_worker_is_rejected_before_claiming_work( + protocol_version: int, monkeypatch: pytest.MonkeyPatch +) -> None: from litellm.proxy.lens.endpoints import claim from tests.unit.proxy.lens.test_state import worker + monkeypatch.setenv("LITELLM_RELEASE_TAG", "v1.2.3") with pytest.raises(HTTPException) as error: - await claim(worker(), protocol_version=1) + await claim(worker(), protocol_version=protocol_version) assert error.value.status_code == 409 assert "Upgrade" in error.value.detail @@ -291,3 +295,38 @@ async def test_preview_reports_calendar_overflow_as_a_validation_error() -> None await preview_sample(body, UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), None) assert error.value.status_code == 422 assert "supported calendar range" in error.value.detail + + +@pytest.mark.asyncio +@pytest.mark.parametrize("worker_release", ("", "v1.2.2", "branch-main-old")) +async def test_different_release_is_rejected_before_accessing_jobs( + monkeypatch: pytest.MonkeyPatch, worker_release: str +) -> None: + from litellm.proxy.lens.endpoints import claim + from litellm.proxy.lens.release import PROTOCOL_VERSION + from tests.unit.proxy.lens.test_state import worker + + monkeypatch.setenv("LITELLM_RELEASE_TAG", "v1.2.3") + monkeypatch.delenv("LENS_WORKER_IMAGE", raising=False) + with pytest.raises(HTTPException) as error: + await claim(worker(), protocol_version=PROTOCOL_VERSION, worker_release=worker_release) + assert error.value.status_code == 409 + assert "ghcr.io/berriai/litellm-lens-worker:v1.2.3" in error.value.detail + + +@pytest.mark.asyncio +async def test_unknown_gateway_release_refuses_registration_and_claims(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy.lens.endpoints import WorkerName, claim, register_worker + from litellm.proxy.lens.release import PROTOCOL_VERSION + from tests.unit.proxy.lens.test_state import worker + + monkeypatch.setenv("LITELLM_RELEASE_TAG", "") + monkeypatch.setenv("LENS_WORKER_IMAGE", "registry.example/lens-worker:old") + with pytest.raises(HTTPException) as registration_error: + await register_worker(WorkerName(analysis_key_id="a" * 64), UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)) + assert registration_error.value.status_code == 503 + assert "LITELLM_RELEASE_TAG" in registration_error.value.detail + with pytest.raises(HTTPException) as claim_error: + await claim(worker(), protocol_version=PROTOCOL_VERSION, worker_release="") + assert claim_error.value.status_code == 503 + assert claim_error.value.detail == registration_error.value.detail diff --git a/tests/unit/proxy/lens/test_release.py b/tests/unit/proxy/lens/test_release.py new file mode 100644 index 00000000000..95d637465c4 --- /dev/null +++ b/tests/unit/proxy/lens/test_release.py @@ -0,0 +1,74 @@ +from importlib.metadata import Distribution, PackageNotFoundError, PathDistribution +from pathlib import Path +from typing import Final + +import pytest + +from litellm.proxy.lens.release import worker_image + + +@pytest.mark.parametrize("tag", ("v1.2.3", "v1.2.3-rc.4", "v1.2.3-dev.5", "branch-main-1234567")) +def test_install_command_follows_the_gateway_release(monkeypatch: pytest.MonkeyPatch, tag: str) -> None: + monkeypatch.setenv("LITELLM_RELEASE_TAG", tag) + monkeypatch.delenv("LENS_WORKER_IMAGE", raising=False) + assert worker_image() == f"ghcr.io/berriai/litellm-lens-worker:{tag}" + + +def test_private_registry_override_keeps_its_exact_digest(monkeypatch: pytest.MonkeyPatch) -> None: + image: Final = "registry.example/lens-worker@sha256:" + "a" * 64 + monkeypatch.setenv("LITELLM_RELEASE_TAG", "branch-main-1234567") + monkeypatch.setenv("LENS_WORKER_IMAGE", image) + assert worker_image() == image + + +@pytest.mark.parametrize( + "installed,expected", + (("1.2.3", "v1.2.3"), ("1.2.3rc4", "v1.2.3-rc.4"), ("1.2.3.dev5", "v1.2.3-dev.5")), +) +def test_python_installs_recommend_the_matching_worker( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, installed: str, expected: str +) -> None: + from litellm.proxy.lens import release + + metadata: Final = tmp_path / "litellm.dist-info" + metadata.mkdir() + metadata.joinpath("METADATA").write_text(f"Name: litellm\nVersion: {installed}\n") + + def installed_distribution(name: str) -> Distribution: + assert name == "litellm" + return PathDistribution(metadata) + + monkeypatch.delenv("LITELLM_RELEASE_TAG", raising=False) + monkeypatch.delenv("LENS_WORKER_IMAGE", raising=False) + monkeypatch.setattr(release, "distribution", installed_distribution) + monkeypatch.setattr(release, "__file__", str(tmp_path / "litellm/proxy/lens/release.py")) + assert release.release_tag() == expected + assert worker_image() == f"ghcr.io/berriai/litellm-lens-worker:{expected}" + + +@pytest.mark.parametrize("source", ("checkout", "direct-install", "unversioned-container", "missing-package")) +def test_unknown_source_never_falls_back_to_a_package_version_or_image_override( + monkeypatch: pytest.MonkeyPatch, tmp_path: Path, source: str +) -> None: + from litellm.proxy.lens import release + + metadata: Final = tmp_path / "litellm.dist-info" + metadata.mkdir() + metadata.joinpath("METADATA").write_text("Name: litellm\nVersion: 1.2.3\n") + if source == "direct-install": + metadata.joinpath("direct_url.json").write_text('{"url":"file:///checkout","dir_info":{"editable":true}}') + + def installed_distribution(name: str) -> Distribution: + if source == "missing-package": + raise PackageNotFoundError(name) + return PathDistribution(metadata) + + monkeypatch.delenv("LITELLM_RELEASE_TAG", raising=False) + monkeypatch.setenv("LENS_WORKER_IMAGE", "registry.example/lens-worker:old") + monkeypatch.setattr(release, "distribution", installed_distribution) + if source != "checkout": + monkeypatch.setattr(release, "__file__", str(tmp_path / "litellm/proxy/lens/release.py")) + if source == "unversioned-container": + monkeypatch.setenv("LITELLM_RELEASE_TAG", "") + assert release.release_tag() == "" + assert worker_image() == "" diff --git a/tests/unit/proxy/lens/test_worker.py b/tests/unit/proxy/lens/test_worker.py index 17bac711e39..7983aec8af4 100644 --- a/tests/unit/proxy/lens/test_worker.py +++ b/tests/unit/proxy/lens/test_worker.py @@ -417,3 +417,22 @@ async def test_transient_heartbeat_failure_recovers_without_cancelling_analysis( assert result.error == "" assert result.coverage.screened == 1 and result.coverage.unassessable == 0 assert attempts.qsize() == 2 and saved.empty() + + +@pytest.mark.asyncio +async def test_worker_announces_release_and_waits_on_incompatible_gateway( + monkeypatch: pytest.MonkeyPatch, caplog: pytest.LogCaptureFixture +) -> None: + from litellm.proxy.lens.release import PROTOCOL_VERSION + + monkeypatch.setenv("LITELLM_RELEASE_TAG", "v1.2.3") + + def handle(request: httpx.Request) -> httpx.Response: + assert request.url.path == "/lens/worker/claim" + assert request.url.params["protocol_version"] == str(PROTOCOL_VERSION) + assert request.url.params["worker_release"] == "v1.2.3" + return httpx.Response(409, json={"detail": "Upgrade the Lens worker to v1.2.4"}) + + async with httpx.AsyncClient(base_url="https://proxy.test", transport=httpx.MockTransport(handle)) as client: + assert not await LensWorker(client).run_once() + assert "Upgrade the Lens worker to v1.2.4" in caplog.text diff --git a/tests/unit/test_lens_dev.py b/tests/unit/test_lens_dev.py index e14b27fd27e..cbc8aced8ac 100644 --- a/tests/unit/test_lens_dev.py +++ b/tests/unit/test_lens_dev.py @@ -2,6 +2,7 @@ import os import subprocess import sys from pathlib import Path +from typing import Final ROOT = Path(__file__).resolve().parents[2] SCRIPT = ROOT / "scripts" / "lens_dev.sh" @@ -123,6 +124,19 @@ def test_proxy_env_permits_the_weak_key_only_when_chosen(tmp_path): assert "LITELLM_DANGEROUSLY_PERMIT_WEAK_OR_UNSET_MASTER_KEY=true" in proc.stdout +def test_source_development_overrides_an_inherited_release_with_its_own_commit(tmp_path: Path) -> None: + proc: Final = _run( + tmp_path, + 'proxy_env "export LITELLM_RELEASE_TAG=v0.0.0-old"; ' + 'test "$LITELLM_RELEASE_TAG" = "sha-$(git -C "$repo_root" rev-parse HEAD)"; ' + 'printf "%s" "$LENS_WORKER_IMAGE"', + LITELLM_RELEASE_TAG="v0.0.0-old", + LENS_WORKER_IMAGE="registry.example/lens-worker:old", + ) + assert proc.returncode == 0, proc.stderr + assert proc.stdout == "litellm-lens-worker:local" + + def test_external_database_url_never_starts_compose_postgres(tmp_path): docker = tmp_path / "bin" / "docker" proc = _run( diff --git a/ui/litellm-dashboard/src/components/lens/setup/worker/WorkerDialog.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/setup/worker/WorkerDialog.integration.test.tsx index 10ee40a87f9..719b83cfe6b 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/worker/WorkerDialog.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/setup/worker/WorkerDialog.integration.test.tsx @@ -12,6 +12,7 @@ vi.mock("@/components/networking", () => ({ const created = { token: "lens-test-token", + image: "ghcr.io/berriai/litellm-lens-worker:v1.2.3", worker: { id: "worker", name: "Lens worker", @@ -57,7 +58,11 @@ describe("Worker setup", () => { expect(command).toContain("LITELLM_URL=https://gateway.example/proxy"); expect(command).toContain("LENS_WORKER_TOKEN=lens-test-token"); expect(command).toContain("--add-host host.docker.internal:host-gateway"); - expect(command).toContain("ghcr.io/berriai/litellm-lens-worker@sha256:"); + expect(command).toContain(created.image); + await user.click(screen.getByText("Using Docker Compose or Helm?")); + await user.click(screen.getByRole("button", { name: "Copy worker token" })); + expect(await navigator.clipboard.readText()).toBe(created.token); + expect(screen.getByRole("button", { name: "Token copied" })).toBeVisible(); }); it("assigns billing to an existing worker without replacing its access token", async () => { const user = userEvent.setup(); diff --git a/ui/litellm-dashboard/src/components/lens/setup/worker/WorkerInstall.tsx b/ui/litellm-dashboard/src/components/lens/setup/worker/WorkerInstall.tsx index 49f4df82edd..efeef194652 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/worker/WorkerInstall.tsx +++ b/ui/litellm-dashboard/src/components/lens/setup/worker/WorkerInstall.tsx @@ -1,5 +1,6 @@ "use client"; +import { useState } from "react"; import { Button } from "@/components/ui/button"; import { CheckCircle2, Copy, Loader2 } from "lucide-react"; @@ -25,6 +26,7 @@ export function WorkerInstall({ onReady?: () => void; onClose: () => void; }) { + const [tokenCopied, setTokenCopied] = useState(false); return (
{!connected && ( @@ -34,7 +36,7 @@ export function WorkerInstall({ className="w-full gap-2" onClick={async () => { try { - await navigator.clipboard.writeText(workerSetupCommand(address, created.token)); + await navigator.clipboard.writeText(workerSetupCommand(address, created.token, created.image)); setCopied(true); } catch { setError("Clipboard access failed. Allow clipboard access and try again."); @@ -51,9 +53,30 @@ export function WorkerInstall({ aria-label="Docker command preview" className="mt-3 max-h-48 overflow-auto rounded-md bg-muted/40 p-3 text-xs leading-5" > - {workerSetupCommand(address, created.token)} + {workerSetupCommand(address, created.token, created.image)} +
+ Using Docker Compose or Helm? +

+ Save this private token as LENS_WORKER_TOKEN in Compose or in your Helm worker token secret. Keep it for + future upgrades. +

+ +
)} {!connected && ( diff --git a/ui/litellm-dashboard/src/components/lens/setup/worker/workerCommand.ts b/ui/litellm-dashboard/src/components/lens/setup/worker/workerCommand.ts index 0b5c1c347d6..0e022aaec26 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/worker/workerCommand.ts +++ b/ui/litellm-dashboard/src/components/lens/setup/worker/workerCommand.ts @@ -1,23 +1,20 @@ import { proxyBaseUrl } from "@/components/networking"; import { serverRootPath } from "@/lib/serverRootPath"; -export const LENS_WORKER_IMAGE = - "ghcr.io/berriai/litellm-lens-worker@sha256:44f0597c7583dcfef999ece9a8bc02cfeb9f0f5167a1221cee3bd10b1b79271b"; - export function initialProxyAddress(): string { const url = new URL(proxyBaseUrl || serverRootPath, window.location.origin); if (["localhost", "127.0.0.1", "[::1]"].includes(url.hostname)) url.hostname = "host.docker.internal"; return url.toString().replace(/\/$/, ""); } -export function workerSetupCommand(address: string, token: string): string { +export function workerSetupCommand(address: string, token: string, image: string): string { const quote = (value: string) => "'" + value.replaceAll("'", "'\\''") + "'"; return [ "docker run -d --restart unless-stopped --read-only --cap-drop ALL", " --tmpfs /tmp:rw,noexec,nosuid,size=1g", - " --security-opt no-new-privileges --platform linux/amd64 --add-host host.docker.internal:host-gateway", + " --security-opt no-new-privileges --add-host host.docker.internal:host-gateway", ` -e ${quote("LITELLM_URL=" + address)}`, ` -e ${quote("LENS_WORKER_TOKEN=" + token)}`, - ` ${LENS_WORKER_IMAGE}`, + ` ${quote(image)}`, ].join(" \\\n"); } diff --git a/ui/litellm-dashboard/src/components/lens/setup/worker/workerSchema.test.ts b/ui/litellm-dashboard/src/components/lens/setup/worker/workerSchema.test.ts index 01e3099260d..fed67e0d497 100644 --- a/ui/litellm-dashboard/src/components/lens/setup/worker/workerSchema.test.ts +++ b/ui/litellm-dashboard/src/components/lens/setup/worker/workerSchema.test.ts @@ -1,6 +1,6 @@ import { describe, expect, it } from "vitest"; import { validateWorkerAddress, analysisAccessSchema, workerFormSchema } from "./workerSchema"; -import { workerSetupCommand, LENS_WORKER_IMAGE } from "./workerCommand"; +import { workerSetupCommand } from "./workerCommand"; const workerDefaults = { useExisting: false, @@ -26,14 +26,13 @@ describe("worker setup", () => { }); it("quotes apostrophes literally and retains the pinned image and runtime restrictions", () => { - const command = workerSetupCommand("https://gateway.example/proxy?name=it's", "token'quoted"); + const image = "registry.example/lens-worker:v1.2.3-rc.4"; + const command = workerSetupCommand("https://gateway.example/proxy?name=it's", "token'quoted", image); expect(command).toContain("'LITELLM_URL=https://gateway.example/proxy?name=it'\\''s'"); expect(command).toContain("'LENS_WORKER_TOKEN=token'\\''quoted'"); expect(command).toContain("--read-only --cap-drop ALL"); - expect(command).toContain( - "--security-opt no-new-privileges --platform linux/amd64 --add-host host.docker.internal:host-gateway", - ); - expect(command.split("\n").at(-1)?.trim()).toBe(LENS_WORKER_IMAGE); + expect(command).toContain("--security-opt no-new-privileges --add-host host.docker.internal:host-gateway"); + expect(command.split("\n").at(-1)?.trim()).toBe(`'${image}'`); }); }); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 5478772bc29..f0423a8215b 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -50091,6 +50091,8 @@ export interface components { }; /** WorkerCreated */ WorkerCreated: { + /** Image */ + image: string; /** Token */ token: string; worker: components["schemas"]["Worker"]; @@ -62899,6 +62901,7 @@ export interface operations { parameters: { query?: { protocol_version?: number; + worker_release?: string; }; header?: never; path?: never; From 1ef0fe9790afc007c54b32c8167faeac54ec5498 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 16:35:11 -0700 Subject: [PATCH 05/18] fix(auth): resolve hidden model_group_alias entries in the zero-cost budget check (#43741) * fix(auth): resolve hidden model_group_alias entries in the zero-cost budget check Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(auth): key the zero-cost cache by the resolved model group Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(auth): keep the zero-cost verdict per requested name and include hidden aliases * test(integration): audit the zero-cost bypass through hidden model_group_alias names * test(router): cover the extracted routing strategy switch * test(integration): record a pre-flip burst before the alias flip --------- Co-authored-by: jesus Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/proxy/auth/auth_checks.py | 2 +- litellm/router.py | 35 +- .../authorization/_hidden_alias_budget.py | 423 +++++++++++++ .../test_hidden_alias_budget_bypass.py | 577 ++++++++++++++++++ .../test_hidden_alias_budget_bypass_chaos.py | 208 +++++++ .../test_unmapped_model_budget_enforcement.py | 150 ++++- tests/unit/test_router/test_router.py | 67 ++ 7 files changed, 1430 insertions(+), 32 deletions(-) create mode 100644 tests/integration/authorization/_hidden_alias_budget.py create mode 100644 tests/integration/authorization/test_hidden_alias_budget_bypass.py create mode 100644 tests/integration/authorization/test_hidden_alias_budget_bypass_chaos.py diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 918428bcf0e..dfdf7dc4e66 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -484,7 +484,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None continue try: # Use router's get_model_group_info method directly for better reliability - model_group_info = llm_router.get_model_group_info(model_group=model_name) + model_group_info = llm_router.get_model_group_info(model_group=model_name, include_hidden=True) if model_group_info is None: # Model not found or no pricing info available diff --git a/litellm/router.py b/litellm/router.py index 49e9b9b5a78..54fffae8dbd 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -11246,13 +11246,13 @@ class Router: return model_group_info - def get_model_group_info(self, model_group: str) -> ModelGroupInfo | None: + def get_model_group_info(self, model_group: str, *, include_hidden: bool = False) -> ModelGroupInfo | None: """ For a given model group name, return the combined model info Returns: - ModelGroupInfo if able to construct a model group - - None if error constructing model group info or hidden model group + - None if error constructing model group info or hidden model group (unless include_hidden) """ ## Check if model group alias if model_group in self.model_group_alias: @@ -11260,7 +11260,7 @@ class Router: if isinstance(item, str): _router_model_group = item elif isinstance(item, dict): - if item["hidden"] is True: + if item["hidden"] is True and not include_hidden: return None else: _router_model_group = item["model"] @@ -12302,6 +12302,16 @@ class Router: ] return _settings_to_return + def _switch_routing_strategy(self, routing_strategy: str | None, kwargs: Mapping[str, object]) -> None: + if routing_strategy == "lar1": + from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy + + apply_lar1_routing_strategy(self, kwargs.get("routing_strategy_args")) + return + self.routing_strategy_init( + routing_strategy=routing_strategy, routing_strategy_args=kwargs.get("routing_strategy_args", {}) + ) + def update_settings(self, **kwargs): """ Update the router settings. @@ -12315,6 +12325,7 @@ class Router: ] _existing_router_settings: Final = self.get_settings() + model_group_alias_before: Final = self.model_group_alias rebuild_routing_groups = False routing_args_updated = False for var in kwargs: @@ -12338,20 +12349,7 @@ class Router: if var == "routing_strategy": value = self._normalize_strategy(value) if _existing_router_settings["routing_strategy"] != value: - if value == "lar1": - from litellm.router_strategy.lar1_routing import ( - apply_lar1_routing_strategy, - ) - - apply_lar1_routing_strategy( - self, - kwargs.get("routing_strategy_args"), - ) - else: - self.routing_strategy_init( - routing_strategy=value, - routing_strategy_args=kwargs.get("routing_strategy_args", {}), - ) + self._switch_routing_strategy(value, kwargs) rebuild_routing_groups = True elif var == "routing_strategy_args": routing_args_updated = value != self.routing_strategy_args @@ -12362,6 +12360,9 @@ class Router: if routing_args_updated: self._apply_updated_routing_strategy_args() + if self.model_group_alias != model_group_alias_before: + self._invalidate_model_group_info_cache() + if rebuild_routing_groups: routing_groups_input: Final = kwargs.get("routing_groups", self._routing_groups_input) self._init_routing_groups(routing_groups_input) diff --git a/tests/integration/authorization/_hidden_alias_budget.py b/tests/integration/authorization/_hidden_alias_budget.py new file mode 100644 index 00000000000..eb099e56026 --- /dev/null +++ b/tests/integration/authorization/_hidden_alias_budget.py @@ -0,0 +1,423 @@ +import json +import os +import uuid +from collections import Counter +from collections.abc import Callable, Generator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from dataclasses import dataclass +from hashlib import sha256 +from types import MappingProxyType +from typing import Final + +import httpx +from anthropic import Anthropic, AsyncAnthropic +from integration._support.anthropic_thinking import JSON_LIST, JSON_OBJECT +from integration._support.client import ( + GATEWAY_LIMITS, + Gateway, + Scenario, + eventually, + gateway_from_environment, + object_value, +) +from integration._support.database import read_rows +from integration._support.upstream import delete_scenario, register_scenario +from integration.cost_calculation.cost_tracking_case import JsonResponse, SseResponse +from openai import AsyncOpenAI, OpenAI +from pydantic import JsonValue + +BUDGET: Final = 0.05 +CHAT_REPLY: Final = "Hello! This is a mock response from the fake OpenAI endpoint." +RESPONSES_REPLY: Final = "free reply" +BUDGET_EXCEEDED: Final = 422 +GATEWAY_BURST: Final = 24 +PEER_BURST: Final = 4 +SPEND_MARKER_HEADER: Final = "x-litellm-spend-logs-metadata" +PROXY_BUDGET_USER: Final = "litellm-proxy-budget" + +_RESPONSE: Final[JsonValue] = { + "id": "resp_$UNIQUE_ID", + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-4o-mini", + "output": [ + { + "id": "msg_$UNIQUE_ID", + "type": "message", + "role": "assistant", + "status": "completed", + "content": [{"type": "output_text", "text": RESPONSES_REPLY, "annotations": []}], + } + ], + "usage": {"input_tokens": 20, "output_tokens": 20, "total_tokens": 40}, +} +_RESPONSE_EVENTS: Final[tuple[Mapping[str, JsonValue], ...]] = ( + { + "type": "response.created", + "sequence_number": 0, + "response": {**_RESPONSE, "status": "in_progress", "output": []}, + }, + { + "type": "response.output_text.delta", + "sequence_number": 1, + "item_id": "msg_$UNIQUE_ID", + "output_index": 0, + "content_index": 0, + "delta": RESPONSES_REPLY, + }, + {"type": "response.completed", "sequence_number": 2, "response": _RESPONSE}, +) + + +@dataclass(frozen=True, slots=True) +class AliasRig: + gateway: Gateway + peer: Gateway + free: str + paid: str + failing_free: str + failing_provider_model: str + hidden_free: str + visible_free: str + hidden_paid: str + hidden_responses: str + hidden_responses_stream: str + hidden_failing: str + shown_free: str + null_hidden_free: str + hidden_unpriced: str + hidden_missing: str + + +def base_url(candidate: Gateway) -> str: + return str(candidate.client.base_url).rstrip("/") + + +def spend_marker(marker: str) -> Mapping[str, str]: + return {SPEND_MARKER_HEADER: json.dumps({"marker": marker})} + + +def fresh_post(candidate: Gateway, path: str, body: Mapping[str, JsonValue], key: str, marker: str) -> httpx.Response: + return httpx.post( + f"{base_url(candidate)}{path}", + json=dict(body), + headers={"Authorization": f"Bearer {key}", **spend_marker(marker)}, + timeout=60, + trust_env=False, + ) + + +NO_EXTRA: Final[Mapping[str, JsonValue]] = MappingProxyType({}) + + +def chat_body(model: JsonValue, marker: str, extra: Mapping[str, JsonValue] = NO_EXTRA) -> Mapping[str, JsonValue]: + return {"model": model, "messages": [{"role": "user", "content": marker}], **extra} + + +def fresh_chat( + candidate: Gateway, model: str, key: str, marker: str, extra: Mapping[str, JsonValue] = NO_EXTRA +) -> httpx.Response: + return fresh_post(candidate, "/v1/chat/completions", chat_body(model, marker, extra), key, marker) + + +def fresh_response( + candidate: Gateway, model: str, key: str, marker: str, extra: Mapping[str, JsonValue] = NO_EXTRA +) -> httpx.Response: + return fresh_post(candidate, "/v1/responses", {"model": model, "input": marker, **extra}, key, marker) + + +def fresh_message( + candidate: Gateway, model: str, key: str, marker: str, extra: Mapping[str, JsonValue] = NO_EXTRA +) -> httpx.Response: + return fresh_post( + candidate, + "/v1/messages", + {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": marker}], **extra}, + key, + marker, + ) + + +def statuses(send: Callable[[str], httpx.Response], count: int) -> frozenset[int]: + markers: Final = tuple(uuid.uuid4().hex for _ in range(count)) + with ThreadPoolExecutor(max_workers=count) as pool: + return frozenset(response.status_code for response in pool.map(send, markers)) + + +def error_type(response: httpx.Response) -> str: + return str(object_value(JSON_OBJECT.validate_json(response.content)["error"])["type"]) + + +def settle( + send: Callable[[str], httpx.Response], status: int, *, burst: int = GATEWAY_BURST, seconds: float = 60 +) -> None: + eventually( + lambda: statuses(send, burst) | statuses(send, burst), + lambda seen: seen == frozenset({status}), + seconds=seconds, + ) + + +def chat_statuses( + candidate: Gateway, model: str, key: str, count: int, extra: Mapping[str, JsonValue] = NO_EXTRA +) -> frozenset[int]: + return statuses(lambda marker: fresh_chat(candidate, model, key, marker, extra), count) + + +def settle_candidate( + candidate: Gateway, + model: str, + key: str, + status: int, + *, + burst: int = GATEWAY_BURST, + seconds: float = 60, + extra: Mapping[str, JsonValue] = NO_EXTRA, +) -> None: + settle(lambda marker: fresh_chat(candidate, model, key, marker, extra), status, burst=burst, seconds=seconds) + + +def settle_chat( + rig: AliasRig, + model: str, + key: str, + status: int, + *, + seconds: float = 60, + extra: Mapping[str, JsonValue] = NO_EXTRA, +) -> None: + settle_candidate(rig.gateway, model, key, status, seconds=seconds, extra=extra) + settle_candidate(rig.peer, model, key, status, burst=PEER_BURST, seconds=seconds, extra=extra) + + +def exhausted_key(rig: AliasRig, scenario: Scenario) -> str: + key: Final = scenario.key(max_budget=BUDGET) + first: Final = fresh_chat(rig.gateway, rig.paid, key, "exhaust-" + uuid.uuid4().hex) + assert first.status_code == 200, first.text + settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED) + return key + + +def upstream_requests(upstream_url: str) -> tuple[str, ...]: + drained: Final = JSON_OBJECT.validate_json( + httpx.get(f"{upstream_url}/__observations", timeout=15, trust_env=False).content + ) + return tuple(json.dumps(entry) for entry in JSON_LIST.validate_python(drained["requests"])) + + +def upstream_hits(observed: tuple[str, ...], marker: str) -> int: + return sum(1 for entry in observed if marker in entry) + + +def script_provider(rig: AliasRig, failures: int) -> None: + scripted: Final = httpx.post( + f"{rig.gateway.upstream_url}/__scripts/{rig.failing_provider_model}", + json={"statuses": [500] * failures}, + timeout=15, + trust_env=False, + ) + assert scripted.status_code == 200, scripted.text + + +def clear_provider_script(rig: AliasRig) -> None: + cleared: Final = httpx.delete( + f"{rig.gateway.upstream_url}/__scripts/{rig.failing_provider_model}", timeout=15, trust_env=False + ) + assert cleared.status_code in (200, 404), cleared.text + + +def spend_rows(key: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple( + read_rows( + "SELECT request_id, spend, model_group, status, call_type, " + "metadata->'spend_logs_metadata'->>'marker' AS marker " + 'FROM "LiteLLM_SpendLogs" WHERE api_key = %s', + (sha256(key.encode()).hexdigest(),), + ) + ) + + +def landed(key: str, marker: str) -> tuple[Mapping[str, JsonValue], ...]: + return tuple(row for row in spend_rows(key) if row["marker"] == marker) + + +def landed_once(key: str, marker: str) -> Mapping[str, JsonValue]: + rows: Final = eventually(lambda: landed(key, marker), lambda found: len(found) >= 1, seconds=70) + assert len(rows) == 1, rows + return rows[0] + + +def marker_counts(key: str, markers: frozenset[str]) -> Mapping[str, int]: + return dict(Counter(str(row["marker"]) for row in spend_rows(key) if row["marker"] in markers)) + + +def landed_all_once(key: str, markers: frozenset[str]) -> tuple[Mapping[str, JsonValue], ...]: + counts: Final = eventually( + lambda: marker_counts(key, markers), lambda found: frozenset(found) == markers, seconds=90 + ) + assert counts == dict.fromkeys(markers, 1), counts + return tuple(row for row in spend_rows(key) if row["marker"] in markers) + + +def assert_free_row(row: Mapping[str, JsonValue], model_group: str) -> None: + assert float(str(row["spend"])) == 0.0, row + assert row["model_group"] == model_group, row + assert row["status"] == "success", row + + +def alias_map(gateway: Gateway) -> Mapping[str, JsonValue]: + current: Final = object_value(gateway.get("/router/settings")["current_values"]).get("model_group_alias") + return object_value(current) if current is not None else {} + + +def write_aliases(gateway: Gateway, aliases: Mapping[str, JsonValue]) -> None: + gateway.post("/config/update", {"router_settings": {"model_group_alias": dict(aliases)}}) + + +def install_aliases(gateway: Gateway, aliases: Mapping[str, JsonValue]) -> None: + write_aliases(gateway, {**alias_map(gateway), **aliases}) + + +def remove_aliases(gateway: Gateway, names: frozenset[str]) -> None: + write_aliases(gateway, {name: target for name, target in alias_map(gateway).items() if name not in names}) + + +def hidden(group: str) -> JsonValue: + return {"model": group, "hidden": True} + + +def openai_client(candidate: Gateway, key: str) -> OpenAI: + return OpenAI( + api_key=key, + base_url=base_url(candidate) + "/v1", + max_retries=0, + http_client=httpx.Client(timeout=60, trust_env=False), + ) + + +def async_openai_client(candidate: Gateway, key: str) -> AsyncOpenAI: + return AsyncOpenAI( + api_key=key, + base_url=base_url(candidate) + "/v1", + max_retries=0, + http_client=httpx.AsyncClient(timeout=60, trust_env=False), + ) + + +def anthropic_client(candidate: Gateway, key: str) -> Anthropic: + return Anthropic( + api_key=key, + base_url=base_url(candidate), + max_retries=0, + http_client=httpx.Client(timeout=60, trust_env=False), + ) + + +def async_anthropic_client(candidate: Gateway, key: str) -> AsyncAnthropic: + return AsyncAnthropic( + api_key=key, + base_url=base_url(candidate), + max_retries=0, + http_client=httpx.AsyncClient(timeout=60, trust_env=False), + ) + + +def _zero_cost_responses_group( + scenario: Scenario, gateway: Gateway, name: str, response: JsonResponse | SseResponse +) -> str: + handle: Final = register_scenario(name, response, control_url=gateway.upstream_url) + scenario.cleanups.callback(delete_scenario, handle) + return scenario.model(api_base=f"{handle.api_base()}/v1", input_cost_per_token=0, output_cost_per_token=0) + + +def settle_responses( + rig: AliasRig, model: str, key: str, status: int, extra: Mapping[str, JsonValue] = NO_EXTRA +) -> None: + settle(lambda marker: fresh_response(rig.gateway, model, key, marker, extra), status) + settle(lambda marker: fresh_response(rig.peer, model, key, marker, extra), status, burst=PEER_BURST) + + +def _await_rig(rig: AliasRig) -> None: + admin: Final = rig.gateway.key + for model in ( + rig.hidden_free, + rig.visible_free, + rig.hidden_paid, + rig.hidden_failing, + rig.shown_free, + rig.null_hidden_free, + rig.hidden_unpriced, + ): + settle_chat(rig, model, admin, 200) + settle_responses(rig, rig.hidden_responses, admin, 200) + settle_responses(rig, rig.hidden_responses_stream, admin, 200, extra={"stream": True}) + + +@contextmanager +def alias_rig() -> Generator[AliasRig]: + suffix: Final = uuid.uuid4().hex[:12] + with ( + gateway_from_environment() as gateway, + httpx.Client( + base_url=os.environ["INTEGRATION_PEER_URL"], timeout=15, trust_env=False, limits=GATEWAY_LIMITS + ) as peer_client, + gateway.scenario() as scenario, + ): + free: Final = scenario.model(input_cost_per_token=0, output_cost_per_token=0) + paid: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002) + unpriced: Final = scenario.model() + failing_provider_model: Final = f"hidden-alias-failing-{suffix}" + failing_free: Final = scenario.model( + model=f"openai/{failing_provider_model}", input_cost_per_token=0, output_cost_per_token=0 + ) + responses_free: Final = _zero_cost_responses_group( + scenario, + gateway, + f"hidden-alias-json-{suffix}", + JsonResponse(content_type="application/json", body=_RESPONSE), + ) + responses_stream_free: Final = _zero_cost_responses_group( + scenario, + gateway, + f"hidden-alias-sse-{suffix}", + SseResponse( + content_type="text/event-stream", + frames=tuple(f"event: {event['type']}\ndata: {json.dumps(event)}" for event in _RESPONSE_EVENTS), + ), + ) + rig: Final = AliasRig( + gateway=gateway, + peer=Gateway(peer_client, gateway.key, gateway.upstream_url), + free=free, + paid=paid, + failing_free=failing_free, + failing_provider_model=failing_provider_model, + hidden_free=f"hidden-free-{suffix}", + visible_free=f"visible-free-{suffix}", + hidden_paid=f"hidden-paid-{suffix}", + hidden_responses=f"hidden-responses-{suffix}", + hidden_responses_stream=f"hidden-responses-stream-{suffix}", + hidden_failing=f"hidden-failing-{suffix}", + shown_free=f"shown-free-{suffix}", + null_hidden_free=f"null-hidden-free-{suffix}", + hidden_unpriced=f"hidden-unpriced-{suffix}", + hidden_missing=f"hidden-missing-{suffix}", + ) + aliases: Final[Mapping[str, JsonValue]] = { + rig.hidden_free: hidden(free), + rig.visible_free: free, + rig.hidden_paid: hidden(paid), + rig.hidden_responses: hidden(responses_free), + rig.hidden_responses_stream: hidden(responses_stream_free), + rig.hidden_failing: hidden(failing_free), + rig.shown_free: {"model": free, "hidden": False}, + rig.null_hidden_free: {"model": free, "hidden": None}, + rig.hidden_unpriced: hidden(unpriced), + rig.hidden_missing: hidden(f"missing-group-{suffix}"), + } + install_aliases(gateway, aliases) + scenario.cleanups.callback(remove_aliases, gateway, frozenset(aliases)) + _await_rig(rig) + yield rig diff --git a/tests/integration/authorization/test_hidden_alias_budget_bypass.py b/tests/integration/authorization/test_hidden_alias_budget_bypass.py new file mode 100644 index 00000000000..46d0852ca8a --- /dev/null +++ b/tests/integration/authorization/test_hidden_alias_budget_bypass.py @@ -0,0 +1,577 @@ +import json +import math +import uuid +from collections.abc import Iterator, Mapping +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from integration._support.anthropic_thinking import JSON_OBJECT +from integration._support.client import Gateway, object_value +from integration._support.process import owned_proxy +from integration.authorization._hidden_alias_budget import ( + BUDGET, + BUDGET_EXCEEDED, + CHAT_REPLY, + PROXY_BUDGET_USER, + RESPONSES_REPLY, + AliasRig, + alias_rig, + anthropic_client, + assert_free_row, + async_anthropic_client, + async_openai_client, + base_url, + chat_statuses, + clear_provider_script, + error_type, + exhausted_key, + fresh_chat, + hidden, + install_aliases, + landed, + landed_once, + openai_client, + remove_aliases, + script_provider, + settle_candidate, + settle_chat, + spend_marker, + upstream_hits, + upstream_requests, +) +from pydantic import JsonValue + +pytestmark: Final = pytest.mark.timeout(240) + + +@pytest.fixture(scope="module") +def rig() -> Iterator[AliasRig]: + with alias_rig() as built: + yield built + + +def test_exhausted_key_reaches_hidden_free_alias_through_openai_chat(rig: AliasRig) -> None: + marker: Final = "chat-sync-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + with openai_client(rig.gateway, key) as client: + completion: Final = client.chat.completions.create( + model=rig.hidden_free, + messages=[{"role": "user", "content": marker}], + extra_headers=spend_marker(marker), + ) + assert completion.choices[0].message.content == CHAT_REPLY, completion + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == completion.id, row + assert_free_row(row, rig.hidden_free) + + +async def test_exhausted_key_reaches_hidden_free_alias_through_streamed_openai_chat(rig: AliasRig) -> None: + marker: Final = "chat-stream-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + async with async_openai_client(rig.gateway, key) as client: + stream: Final = await client.chat.completions.create( + model=rig.hidden_free, + messages=[{"role": "user", "content": marker}], + stream=True, + stream_options={"include_usage": True}, + extra_headers=spend_marker(marker), + ) + chunks: Final = tuple([chunk async for chunk in stream]) + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks if chunk.choices) == CHAT_REPLY + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == chunks[0].id, row + assert_free_row(row, rig.hidden_free) + + +def test_exhausted_key_reaches_hidden_free_alias_through_anthropic_messages(rig: AliasRig) -> None: + marker: Final = "messages-sync-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + with anthropic_client(rig.gateway, key) as client: + message: Final = client.messages.create( + model=rig.hidden_responses, + max_tokens=16, + messages=[{"role": "user", "content": marker}], + extra_headers=spend_marker(marker), + ) + assert [block.text for block in message.content if block.type == "text"] == [RESPONSES_REPLY], message + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == message.id, row + assert_free_row(row, rig.hidden_responses) + + +async def test_exhausted_key_reaches_hidden_free_alias_through_streamed_anthropic_messages(rig: AliasRig) -> None: + marker: Final = "messages-stream-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + async with ( + async_anthropic_client(rig.gateway, key) as client, + client.messages.stream( + model=rig.hidden_responses_stream, + max_tokens=16, + messages=[{"role": "user", "content": marker}], + extra_headers=spend_marker(marker), + ) as stream, + ): + text: Final = "".join([piece async for piece in stream.text_stream]) + final: Final = await stream.get_final_message() + assert text == RESPONSES_REPLY, final + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == final.id, row + assert_free_row(row, rig.hidden_responses_stream) + + +def test_exhausted_key_reaches_hidden_free_alias_through_openai_responses(rig: AliasRig) -> None: + marker: Final = "responses-sync-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + with openai_client(rig.gateway, key) as client: + response: Final = client.responses.create( + model=rig.hidden_responses, input=marker, extra_headers=spend_marker(marker) + ) + assert response.output_text == RESPONSES_REPLY, response + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == response.id, row + assert_free_row(row, rig.hidden_responses) + + +async def test_exhausted_key_reaches_hidden_free_alias_through_streamed_openai_responses(rig: AliasRig) -> None: + marker: Final = "responses-stream-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + async with async_openai_client(rig.gateway, key) as client: + stream: Final = await client.responses.create( + model=rig.hidden_responses_stream, input=marker, stream=True, extra_headers=spend_marker(marker) + ) + events: Final = tuple([event async for event in stream]) + assert events[-1].type == "response.completed", events + assert "".join(event.delta for event in events if event.type == "response.output_text.delta") == RESPONSES_REPLY + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + assert_free_row(landed_once(key, marker), rig.hidden_responses_stream) + + +def _assert_raw_chat_served(rig: AliasRig, candidate: Gateway, key: str) -> None: + marker: Final = "raw-" + uuid.uuid4().hex + response: Final = fresh_chat(candidate, rig.hidden_free, key, marker) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == CHAT_REPLY, response.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + row: Final = landed_once(key, marker) + assert row["request_id"] == response.json()["id"], row + assert_free_row(row, rig.hidden_free) + + +def test_exhausted_key_reaches_hidden_free_alias_over_raw_http_on_both_replicas(rig: AliasRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + _assert_raw_chat_served(rig, rig.gateway, key) + _assert_raw_chat_served(rig, rig.peer, key) + + +def _duplicate_model_post(rig: AliasRig, key: str, first: str, last: str, marker: str) -> httpx.Response: + messages: Final = json.dumps([{"role": "user", "content": marker}]) + return httpx.post( + f"{base_url(rig.gateway)}/v1/chat/completions", + content=f'{{"model": {json.dumps(first)}, "model": {json.dumps(last)}, "messages": {messages}}}'.encode(), + headers={"Authorization": f"Bearer {key}", "Content-Type": "application/json", **spend_marker(marker)}, + timeout=60, + trust_env=False, + ) + + +def test_duplicate_model_field_is_judged_by_its_last_value(rig: AliasRig) -> None: + free_marker: Final = "duplicate-free-" + uuid.uuid4().hex + paid_marker: Final = "duplicate-paid-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + served: Final = _duplicate_model_post(rig, key, rig.hidden_paid, rig.hidden_free, free_marker) + assert served.status_code == 200, served.text + refused: Final = _duplicate_model_post(rig, key, rig.hidden_free, rig.hidden_paid, paid_marker) + assert refused.status_code == BUDGET_EXCEEDED, refused.text + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert upstream_hits(observed, free_marker) == 1 + assert upstream_hits(observed, paid_marker) == 0 + assert_free_row(landed_once(key, free_marker), rig.hidden_free) + + +def test_provider_failure_behind_hidden_free_alias_reaches_the_caller(rig: AliasRig) -> None: + failed_marker: Final = "provider-failure-" + uuid.uuid4().hex + recovered_marker: Final = "provider-recovered-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + scenario.cleanups.callback(clear_provider_script, rig) + script_provider(rig, 1) + failed: Final = fresh_chat(rig.gateway, rig.hidden_failing, key, failed_marker) + assert failed.status_code == 500, failed.text + assert "Controlled provider failure" in failed.text, failed.text + assert "budget" not in failed.text.lower(), failed.text + clear_provider_script(rig) + recovered: Final = fresh_chat(rig.gateway, rig.hidden_failing, key, recovered_marker) + assert recovered.status_code == 200, recovered.text + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert upstream_hits(observed, failed_marker) == 1 + assert upstream_hits(observed, recovered_marker) == 1 + assert_free_row(landed_once(key, recovered_marker), rig.hidden_failing) + + +def test_exhausted_user_budget_still_reaches_hidden_free_alias(rig: AliasRig) -> None: + with rig.gateway.scenario() as scenario: + user: Final = scenario.user(max_budget=BUDGET) + key: Final = scenario.key(user_id=user, max_budget=5.0) + first: Final = fresh_chat(rig.gateway, rig.paid, key, "user-exhaust-" + uuid.uuid4().hex) + assert first.status_code == 200, first.text + settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED, seconds=90) + assert chat_statuses(rig.gateway, rig.hidden_free, key, 8) == {200} + assert chat_statuses(rig.peer, rig.hidden_free, key, 8) == {200} + settle_chat(rig, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=90) + refused: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "user-refused-" + uuid.uuid4().hex) + assert f"User={user}" in refused.text, refused.text + + +def test_exhausted_team_budget_still_reaches_hidden_free_alias(rig: AliasRig) -> None: + with rig.gateway.scenario() as scenario: + team: Final = scenario.team(max_budget=BUDGET) + key: Final = scenario.key(team_id=team, max_budget=5.0) + first: Final = fresh_chat(rig.gateway, rig.paid, key, "team-exhaust-" + uuid.uuid4().hex) + assert first.status_code == 200, first.text + settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED, seconds=90) + assert chat_statuses(rig.gateway, rig.hidden_free, key, 8) == {200} + assert chat_statuses(rig.peer, rig.hidden_free, key, 8) == {200} + settle_chat(rig, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=90) + refused: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "team-refused-" + uuid.uuid4().hex) + assert f"Team={team}" in refused.text, refused.text + + +def test_exhausted_tag_budget_still_reaches_hidden_free_alias(rig: AliasRig) -> None: + tag: Final = "hidden-alias-tag-" + uuid.uuid4().hex + tagged: Final[Mapping[str, JsonValue]] = {"metadata": {"tags": [tag]}} + with rig.gateway.scenario() as scenario: + rig.gateway.post("/tag/new", {"name": tag, "max_budget": BUDGET}) + scenario.cleanups.callback(rig.gateway.post, "/tag/delete", {"name": tag}) + key: Final = scenario.key(max_budget=5.0) + first: Final = fresh_chat(rig.gateway, rig.paid, key, "tag-exhaust-" + uuid.uuid4().hex, tagged) + assert first.status_code == 200, first.text + settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED, seconds=90, extra=tagged) + assert chat_statuses(rig.gateway, rig.hidden_free, key, 8, tagged) == {200} + assert chat_statuses(rig.peer, rig.hidden_free, key, 8, tagged) == {200} + settle_chat(rig, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=90, extra=tagged) + refused: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "tag-refused-" + uuid.uuid4().hex, tagged) + assert f"Tag={tag}" in refused.text, refused.text + untagged: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "tag-untagged-" + uuid.uuid4().hex) + assert untagged.status_code == 200, untagged.text + + +def test_hidden_alias_repointed_between_paid_and_free_groups_follows_the_target(rig: AliasRig) -> None: + alias: Final = "hidden-repoint-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + install_aliases(rig.gateway, {alias: hidden(rig.paid)}) + scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({alias})) + settle_chat(rig, alias, rig.gateway.key, 200) + key: Final = exhausted_key(rig, scenario) + settle_chat(rig, alias, key, BUDGET_EXCEEDED) + install_aliases(rig.gateway, {alias: hidden(rig.free)}) + settle_chat(rig, alias, key, 200) + install_aliases(rig.gateway, {alias: hidden(rig.paid)}) + settle_chat(rig, alias, key, BUDGET_EXCEEDED) + + +def test_visible_alias_repointed_to_a_paid_group_loses_the_bypass(rig: AliasRig) -> None: + alias: Final = "visible-repoint-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + install_aliases(rig.gateway, {alias: rig.free}) + scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({alias})) + key: Final = exhausted_key(rig, scenario) + settle_chat(rig, alias, key, 200) + install_aliases(rig.gateway, {alias: rig.paid}) + settle_chat(rig, alias, key, BUDGET_EXCEEDED) + install_aliases(rig.gateway, {alias: rig.free}) + settle_chat(rig, alias, key, 200) + + +def test_failed_free_primary_falls_back_to_hidden_free_alias_for_exhausted_key(rig: AliasRig) -> None: + free_marker: Final = "fallback-free-" + uuid.uuid4().hex + paid_marker: Final = "fallback-paid-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + scenario.cleanups.callback(clear_provider_script, rig) + script_provider(rig, 2) + served: Final = fresh_chat(rig.gateway, rig.failing_free, key, free_marker, {"fallbacks": [rig.hidden_free]}) + assert served.status_code == 200, served.text + assert served.headers["x-litellm-model-group"] == rig.hidden_free, dict(served.headers) + refused: Final = fresh_chat(rig.gateway, rig.failing_free, key, paid_marker, {"fallbacks": [rig.hidden_paid]}) + assert refused.status_code == 500, refused.text + assert "Controlled provider failure" in refused.text, refused.text + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert upstream_hits(observed, free_marker) == 2 + assert upstream_hits(observed, paid_marker) == 1 + assert_free_row(landed_once(key, free_marker), rig.hidden_free) + + +def test_exhausted_key_is_served_a_cached_reply_through_hidden_free_alias(rig: AliasRig) -> None: + marker: Final = "cache-twin-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + first: Final = fresh_chat(rig.gateway, rig.hidden_free, key, marker) + assert first.status_code == 200, first.text + second: Final = fresh_chat(rig.gateway, rig.hidden_free, key, marker) + assert second.status_code == 200, second.text + assert second.json()["id"] == first.json()["id"], second.text + assert "x-litellm-cache-key" in second.headers, dict(second.headers) + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 1 + rows: Final = landed(key, marker) + assert all(float(str(row["spend"])) == 0.0 for row in rows), rows + + +def _base_config() -> Mapping[str, JsonValue]: + return JSON_OBJECT.validate_python(yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text())) + + +def _own_config(directory: Path, name: str, section: str, value: JsonValue) -> Path: + path: Final = directory / name + path.write_text(yaml.safe_dump({**_base_config(), section: value})) + return path + + +def _delete_proxy_budget_row(rig: AliasRig) -> None: + deleted: Final = rig.gateway.request("POST", "/user/delete", {"user_ids": [PROXY_BUDGET_USER]}) + assert deleted.status_code == 200, deleted.text + + +@pytest.mark.timeout(480) +def test_exhausted_proxy_budget_still_reaches_hidden_free_alias(rig: AliasRig, tmp_path: Path) -> None: + settings: Final = object_value(_base_config()["litellm_settings"]) + config: Final = _own_config( + tmp_path, + "proxy-budget.yaml", + "litellm_settings", + {**settings, "max_budget": BUDGET, "budget_duration": "30d"}, + ) + with rig.gateway.scenario() as scenario: + key: Final = scenario.key() + scenario.cleanups.callback(_delete_proxy_budget_row, rig) + with owned_proxy(rig.gateway, tmp_path, {}, config=config, workers=2) as candidate: + settle_candidate(candidate, rig.hidden_free, key, 200, seconds=120) + first: Final = fresh_chat(candidate, rig.paid, key, "proxy-exhaust-" + uuid.uuid4().hex) + assert first.status_code == 200, first.text + settle_candidate(candidate, rig.paid, key, BUDGET_EXCEEDED, seconds=120) + refused: Final = fresh_chat(candidate, rig.paid, key, "proxy-refused-" + uuid.uuid4().hex) + assert error_type(refused) == "budget_exceeded", refused.text + assert "Key=" not in refused.text, refused.text + assert chat_statuses(candidate, rig.hidden_free, key, 8) == {200} + settle_candidate(candidate, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=120) + unbudgeted: Final = fresh_chat(rig.gateway, rig.paid, key, "proxy-unbudgeted-" + uuid.uuid4().hex) + assert unbudgeted.status_code == 200, unbudgeted.text + + +_TAG_ADDER: Final = """from litellm.integrations.custom_guardrail import CustomGuardrail + + +class TagAdder(CustomGuardrail): + async def async_pre_call_hook(self, user_api_key_dict, cache, data, call_type): + metadata = data.setdefault("metadata", {}) + metadata["tags"] = [*(metadata.get("tags") or []), "__TAG__"] + return data +""" + + +@pytest.mark.timeout(480) +def test_guardrail_added_tag_over_budget_still_reaches_hidden_free_alias(rig: AliasRig, tmp_path: Path) -> None: + tag: Final = "hidden-alias-guardrail-tag-" + uuid.uuid4().hex + module: Final = "tag_adder_" + uuid.uuid4().hex + served_marker: Final = "guardrail-served-" + uuid.uuid4().hex + (tmp_path / f"{module}.py").write_text(_TAG_ADDER.replace("__TAG__", tag)) + config: Final = _own_config( + tmp_path, + "guardrail-tag.yaml", + "guardrails", + [ + { + "guardrail_name": "tag-adder-" + uuid.uuid4().hex, + "litellm_params": {"guardrail": f"{module}.TagAdder", "mode": "pre_call", "default_on": True}, + } + ], + ) + tagged: Final[Mapping[str, JsonValue]] = {"metadata": {"tags": [tag]}} + with rig.gateway.scenario() as scenario: + rig.gateway.post("/tag/new", {"name": tag, "max_budget": BUDGET}) + scenario.cleanups.callback(rig.gateway.post, "/tag/delete", {"name": tag}) + key: Final = scenario.key(max_budget=5.0) + first: Final = fresh_chat(rig.gateway, rig.paid, key, "guardrail-exhaust-" + uuid.uuid4().hex, tagged) + assert first.status_code == 200, first.text + settle_chat(rig, rig.paid, key, BUDGET_EXCEEDED, seconds=90, extra=tagged) + untagged: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "guardrail-untagged-" + uuid.uuid4().hex) + assert untagged.status_code == 200, untagged.text + with owned_proxy(rig.gateway, tmp_path, {}, config=config, workers=2) as candidate: + settle_candidate(candidate, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=120) + refused: Final = fresh_chat(candidate, rig.hidden_paid, key, "guardrail-refused-" + uuid.uuid4().hex) + assert f"Tag={tag}" in refused.text, refused.text + assert chat_statuses(candidate, rig.hidden_free, key, 8) == {200} + served: Final = fresh_chat(candidate, rig.hidden_free, key, served_marker) + assert served.status_code == 200, served.text + row: Final = landed_once(key, served_marker) + assert row["request_id"] == served.json()["id"], row + assert_free_row(row, rig.hidden_free) + + +def _assert_free_alias_served(rig: AliasRig, candidate: Gateway, alias: str, key: str, prefix: str) -> None: + marker: Final = f"{prefix}-" + uuid.uuid4().hex + response: Final = fresh_chat(candidate, alias, key, marker) + assert response.status_code == 200, response.text + assert_free_row(landed_once(key, marker), alias) + + +def test_exhausted_key_reaches_visible_free_alias(rig: AliasRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + _assert_free_alias_served(rig, rig.gateway, rig.visible_free, key, "visible") + _assert_free_alias_served(rig, rig.peer, rig.visible_free, key, "visible") + + +def _assert_hidden_paid_refused(rig: AliasRig, candidate: Gateway, key: str) -> None: + marker: Final = "hidden-paid-" + uuid.uuid4().hex + response: Final = fresh_chat(candidate, rig.hidden_paid, key, marker) + assert response.status_code == BUDGET_EXCEEDED, response.text + assert error_type(response) == "budget_exceeded", response.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0 + + +def test_exhausted_key_is_refused_on_hidden_paid_alias(rig: AliasRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + _assert_hidden_paid_refused(rig, rig.gateway, key) + _assert_hidden_paid_refused(rig, rig.peer, key) + + +def test_exhausted_key_reaches_free_group_by_its_own_name(rig: AliasRig) -> None: + marker: Final = "plain-free-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + response: Final = fresh_chat(rig.gateway, rig.free, key, marker) + assert response.status_code == 200, response.text + assert_free_row(landed_once(key, marker), rig.free) + + +def test_key_with_headroom_is_billed_through_hidden_paid_alias(rig: AliasRig) -> None: + paid_marker: Final = "headroom-paid-" + uuid.uuid4().hex + free_marker: Final = "headroom-free-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = scenario.key(max_budget=5.0) + paid: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, paid_marker) + assert paid.status_code == 200, paid.text + free: Final = fresh_chat(rig.gateway, rig.hidden_free, key, free_marker) + assert free.status_code == 200, free.text + billed: Final = landed_once(key, paid_marker) + assert math.isclose(float(str(billed["spend"])), 20 * 0.001 + 20 * 0.002), billed + assert billed["model_group"] == rig.hidden_paid, billed + assert_free_row(landed_once(key, free_marker), rig.hidden_free) + + +def test_key_restricted_to_the_free_group_reaches_its_hidden_alias(rig: AliasRig) -> None: + marker: Final = "restricted-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = scenario.key(models=[rig.free], max_budget=5.0) + response: Final = fresh_chat(rig.gateway, rig.hidden_free, key, marker) + assert response.status_code == 200, response.text + refused: Final = fresh_chat(rig.gateway, rig.hidden_paid, key, "restricted-paid-" + uuid.uuid4().hex) + assert refused.status_code == 403, refused.text + assert error_type(refused) == "key_model_access_denied", refused.text + assert_free_row(landed_once(key, marker), rig.hidden_free) + + +@pytest.mark.parametrize("flag", ["false", "null"]) +def test_alias_with_a_non_hidden_flag_keeps_the_bypass(rig: AliasRig, flag: str) -> None: + alias: Final = {"false": rig.shown_free, "null": rig.null_hidden_free}[flag] + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + _assert_free_alias_served(rig, rig.gateway, alias, key, f"flag-{flag}") + _assert_free_alias_served(rig, rig.peer, alias, key, f"flag-{flag}") + + +def test_hidden_alias_to_a_group_priced_by_the_cost_map_stays_budgeted(rig: AliasRig) -> None: + marker: Final = "unpriced-" + uuid.uuid4().hex + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + response: Final = fresh_chat(rig.gateway, rig.hidden_unpriced, key, marker) + assert response.status_code == BUDGET_EXCEEDED, response.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0 + + +def test_hidden_alias_to_a_missing_group_is_refused_and_the_proxy_stays_healthy(rig: AliasRig) -> None: + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + refused: Final = fresh_chat(rig.gateway, rig.hidden_missing, key, "missing-" + uuid.uuid4().hex) + assert refused.status_code == BUDGET_EXCEEDED, refused.text + unroutable: Final = fresh_chat(rig.gateway, rig.hidden_missing, rig.gateway.key, "missing-" + uuid.uuid4().hex) + assert unroutable.status_code == 400, unroutable.text + assert "no healthy deployments" in unroutable.text, unroutable.text + for candidate in (rig.gateway, rig.peer): + assert candidate.request("GET", "/health/liveliness").status_code == 200 + assert candidate.request("GET", "/model/info").status_code == 200 + assert candidate.request("GET", "/v1/models").status_code == 200 + served: Final = fresh_chat(rig.gateway, rig.hidden_free, key, "missing-control-" + uuid.uuid4().hex) + assert served.status_code == 200, served.text + + +@pytest.mark.parametrize( + ("shape", "status"), + [("int", BUDGET_EXCEEDED), ("list", 400), ("empty", BUDGET_EXCEEDED), ("oversized", BUDGET_EXCEEDED)], +) +def test_malformed_model_value_never_takes_the_bypass(rig: AliasRig, shape: str, status: int) -> None: + marker: Final = f"malformed-{shape}-" + uuid.uuid4().hex + models: Final[Mapping[str, JsonValue]] = { + "int": 5, + "list": [rig.hidden_free], + "empty": "", + "oversized": rig.hidden_free + "x" * 5120, + } + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + response: Final = rig.gateway.request( + "POST", + "/v1/chat/completions", + {"model": models[shape], "messages": [{"role": "user", "content": marker}]}, + key=key, + ) + assert response.status_code == status, response.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0 + served: Final = fresh_chat(rig.gateway, rig.hidden_free, key, "malformed-control-" + uuid.uuid4().hex) + assert served.status_code == 200, served.text + + +def test_unauthenticated_request_to_hidden_alias_is_rejected(rig: AliasRig) -> None: + marker: Final = "unauthenticated-" + uuid.uuid4().hex + response: Final = httpx.post( + f"{base_url(rig.gateway)}/v1/chat/completions", + json={"model": rig.hidden_free, "messages": [{"role": "user", "content": marker}]}, + timeout=60, + trust_env=False, + ) + assert response.status_code == 401, response.text + assert upstream_hits(upstream_requests(rig.gateway.upstream_url), marker) == 0 + + +def _assert_hidden_aliases_unlisted(rig: AliasRig, candidate: Gateway) -> None: + models: Final = candidate.get("/v1/models")["data"] + groups: Final = candidate.get("/model_group/info")["data"] + assert isinstance(models, list) and isinstance(groups, list) + listed: Final = frozenset(str(object_value(entry)["id"]) for entry in models) + described: Final = frozenset(str(object_value(entry)["model_group"]) for entry in groups) + assert rig.visible_free in listed and rig.visible_free in described + assert rig.shown_free in listed and rig.shown_free in described + for name in (rig.hidden_free, rig.hidden_paid, rig.hidden_responses, rig.hidden_missing): + assert name not in listed and name not in described, name + + +def test_hidden_alias_stays_out_of_model_listings(rig: AliasRig) -> None: + _assert_hidden_aliases_unlisted(rig, rig.gateway) + _assert_hidden_aliases_unlisted(rig, rig.peer) diff --git a/tests/integration/authorization/test_hidden_alias_budget_bypass_chaos.py b/tests/integration/authorization/test_hidden_alias_budget_bypass_chaos.py new file mode 100644 index 00000000000..a7da04f1364 --- /dev/null +++ b/tests/integration/authorization/test_hidden_alias_budget_bypass_chaos.py @@ -0,0 +1,208 @@ +import os +import re +import signal +import threading +import uuid +from collections.abc import Callable, Iterator, Mapping +from concurrent.futures import ThreadPoolExecutor +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +import pytest +from pydantic import JsonValue +from integration._support.client import Gateway, eventually +from integration._support.process import owned_proxy_process, owned_upstream +from integration.authorization._hidden_alias_budget import ( + BUDGET_EXCEEDED, + AliasRig, + alias_rig, + assert_free_row, + chat_statuses, + exhausted_key, + fresh_chat, + fresh_message, + fresh_response, + hidden, + install_aliases, + landed_all_once, + remove_aliases, + settle_candidate, + settle_chat, + upstream_hits, + upstream_requests, +) + +pytestmark: Final = pytest.mark.timeout(240) + +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_WAVE: Final = 8 +_STREAM: Final[Mapping[str, JsonValue]] = {"stream": True} + + +@pytest.fixture(scope="module") +def rig() -> Iterator[AliasRig]: + with alias_rig() as built: + yield built + + +def _burst_markers(prefix: str, count: int) -> tuple[str, ...]: + return tuple(f"{prefix}-{index}-{uuid.uuid4().hex}" for index in range(count)) + + +def _send_all(send: Callable[[str], httpx.Response], markers: tuple[str, ...]) -> tuple[httpx.Response, ...]: + with ThreadPoolExecutor(max_workers=len(markers)) as pool: + return tuple(pool.map(send, markers)) + + +def _mixed_call(rig: AliasRig, key: str, marker: str) -> httpx.Response: + senders: Final[Mapping[str, Callable[[], httpx.Response]]] = { + "chat": lambda: fresh_chat(rig.gateway, rig.hidden_free, key, marker), + "chatstream": lambda: fresh_chat(rig.gateway, rig.hidden_free, key, marker, _STREAM), + "messages": lambda: fresh_message(rig.gateway, rig.hidden_responses, key, marker), + "messagesstream": lambda: fresh_message(rig.gateway, rig.hidden_responses_stream, key, marker, _STREAM), + "responses": lambda: fresh_response(rig.gateway, rig.hidden_responses, key, marker), + "responsesstream": lambda: fresh_response(rig.gateway, rig.hidden_responses_stream, key, marker, _STREAM), + } + return senders[marker.split("-")[1]]() + + +_KINDS: Final = ("chat", "chatstream", "messages", "messagesstream", "responses", "responsesstream") + + +def test_mixed_concurrent_burst_through_hidden_free_aliases_lands_each_call_once(rig: AliasRig) -> None: + markers: Final = tuple(f"burst-{_KINDS[index % len(_KINDS)]}-{index}-{uuid.uuid4().hex}" for index in range(30)) + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + upstream_requests(rig.gateway.upstream_url) + responses: Final = _send_all(lambda marker: _mixed_call(rig, key, marker), markers) + assert [response.status_code for response in responses] == [200] * len(markers), [ + response.text for response in responses if response.status_code != 200 + ] + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert {marker: upstream_hits(observed, marker) for marker in markers} == dict.fromkeys(markers, 1) + rows: Final = landed_all_once(key, frozenset(markers)) + assert len(rows) == len(markers), rows + for row in rows: + assert float(str(row["spend"])) == 0.0, row + assert row["status"] == "success", row + + +def test_alias_flipped_to_a_free_group_during_a_burst_only_ever_serves_or_refuses(rig: AliasRig) -> None: + alias: Final = "flip-" + uuid.uuid4().hex + seen: Final[SimpleQueue[tuple[str, int]]] = SimpleQueue() + with rig.gateway.scenario() as scenario: + install_aliases(rig.gateway, {alias: hidden(rig.paid)}) + scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({alias})) + settle_chat(rig, alias, rig.gateway.key, 200) + key: Final = exhausted_key(rig, scenario) + settle_chat(rig, alias, key, BUDGET_EXCEEDED) + upstream_requests(rig.gateway.upstream_url) + + def wave(candidate: Gateway) -> frozenset[int]: + markers: Final = _burst_markers("flip", _WAVE) + responses: Final = _send_all(lambda marker: fresh_chat(candidate, alias, key, marker), markers) + for marker, response in zip(markers, responses, strict=True): + seen.put((marker, response.status_code)) + return frozenset(response.status_code for response in responses) + + assert wave(rig.gateway) == frozenset({BUDGET_EXCEEDED}) + flip: Final = threading.Thread(target=install_aliases, args=(rig.gateway, {alias: hidden(rig.free)})) + flip.start() + eventually(lambda: wave(rig.gateway) | wave(rig.gateway), lambda found: found == frozenset({200}), seconds=60) + flip.join(timeout=30) + assert not flip.is_alive() + settle_chat(rig, alias, key, 200) + collected: Final = tuple(seen.get() for _ in range(seen.qsize())) + assert {status for _, status in collected} <= {200, BUDGET_EXCEEDED}, collected + served: Final = frozenset(marker for marker, status in collected if status == 200) + refused: Final = frozenset(marker for marker, status in collected if status == BUDGET_EXCEEDED) + assert served and refused, collected + observed: Final = upstream_requests(rig.gateway.upstream_url) + assert all(upstream_hits(observed, marker) == 1 for marker in served), collected + assert all(upstream_hits(observed, marker) == 0 for marker in refused), collected + for row in landed_all_once(key, served): + assert_free_row(row, alias) + + +def _tolerant_status(candidate: Gateway, model: str, key: str, marker: str) -> int | None: + try: + return fresh_chat(candidate, model, key, marker).status_code + except httpx.TransportError: + return None + + +@pytest.mark.timeout(480) +def test_upstream_outage_behind_hidden_free_alias_is_a_provider_error_and_recovers( + rig: AliasRig, tmp_path: Path +) -> None: + alias: Final = "hidden-outage-" + uuid.uuid4().hex + with owned_upstream(tmp_path) as slot, rig.gateway.scenario() as scenario: + group: Final = scenario.model(api_base=f"{slot.url}/v1", input_cost_per_token=0, output_cost_per_token=0) + install_aliases(rig.gateway, {alias: hidden(group)}) + scenario.cleanups.callback(remove_aliases, rig.gateway, frozenset({alias})) + settle_chat(rig, alias, rig.gateway.key, 200) + key: Final = exhausted_key(rig, scenario) + before: Final = _burst_markers("outage-before", 10) + served_before: Final = _send_all(lambda marker: fresh_chat(rig.gateway, alias, key, marker), before) + assert [response.status_code for response in served_before] == [200] * 10 + slot.stop() + during: Final = _burst_markers("outage-during", 10) + failed: Final = _send_all(lambda marker: fresh_chat(rig.gateway, alias, key, marker), during) + for response in failed: + assert response.status_code >= 500, response.text + assert "budget" not in response.text.lower(), response.text + assert rig.gateway.request("GET", "/health/liveliness").status_code == 200 + unrelated: Final = fresh_chat(rig.gateway, rig.hidden_free, key, "outage-unrelated-" + uuid.uuid4().hex) + assert unrelated.status_code == 200, unrelated.text + slot.start() + settle_candidate(rig.gateway, alias, key, 200, seconds=90) + after: Final = _burst_markers("outage-after", 10) + served_after: Final = _send_all(lambda marker: fresh_chat(rig.gateway, alias, key, marker), after) + assert [response.status_code for response in served_after] == [200] * 10 + observed: Final = upstream_requests(slot.url) + assert {marker: upstream_hits(observed, marker) for marker in after} == dict.fromkeys(after, 1) + for row in landed_all_once(key, frozenset(before + after)): + assert_free_row(row, alias) + + +def _worker_startups(log: Path) -> tuple[tuple[int, ...], int]: + text: Final = log.read_text() + started: Final = tuple(int(found[1]) for found in _STARTED_WORKER.finditer(text)) + return started, text.count("Application startup complete.") + + +@pytest.mark.timeout(480) +def test_killed_worker_leaves_the_sibling_serving_hidden_free_aliases(rig: AliasRig, tmp_path: Path) -> None: + with rig.gateway.scenario() as scenario: + key: Final = exhausted_key(rig, scenario) + with owned_proxy_process(rig.gateway, tmp_path, {}, workers=2) as owned: + candidate: Final = owned.gateway + workers, _ = eventually( + lambda: _worker_startups(owned.log), + lambda found: len(found[0]) == 2 and found[1] == 2, + seconds=120, + ) + settle_candidate(candidate, rig.hidden_free, key, 200, seconds=120) + settle_candidate(candidate, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=120) + os.kill(workers[0], signal.SIGKILL) + eventually( + lambda: _tolerant_status(candidate, rig.hidden_free, key, "kill-probe-" + uuid.uuid4().hex), + lambda found: found == 200, + seconds=60, + ) + assert chat_statuses(candidate, rig.hidden_free, key, 8) == {200} + assert chat_statuses(candidate, rig.hidden_paid, key, 8) == {BUDGET_EXCEEDED} + eventually( + lambda: _worker_startups(owned.log), + lambda found: len(found[0]) == 3 and found[1] == 3, + seconds=180, + ) + settle_candidate(candidate, rig.hidden_free, key, 200, seconds=120) + settle_candidate(candidate, rig.hidden_paid, key, BUDGET_EXCEEDED, seconds=120) + markers: Final = _burst_markers("kill-after", 10) + served: Final = _send_all(lambda marker: fresh_chat(candidate, rig.hidden_free, key, marker), markers) + assert [response.status_code for response in served] == [200] * 10 + for row in landed_all_once(key, frozenset(markers)): + assert_free_row(row, rig.hidden_free) diff --git a/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py b/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py index 7665008a6a6..515b6e45dc4 100644 --- a/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py +++ b/tests/unit/proxy/auth/test_unmapped_model_budget_enforcement.py @@ -9,6 +9,8 @@ See: https://github.com/BerriAI/litellm/issues/24770 import copy +import pytest + import litellm from litellm.proxy.auth.auth_checks import _is_model_cost_zero from litellm.router import Router @@ -39,10 +41,7 @@ class TestUnmappedModelBudgetEnforcement: ] ) result = _is_model_cost_zero(model="custom-model", llm_router=router) - assert result is False, ( - "Unmapped model should enforce budget (return False), " - "not bypass it (return True)" - ) + assert result is False, "Unmapped model should enforce budget (return False), not bypass it (return True)" def test_explicitly_free_model_bypasses_budget(self): """A model with explicit cost=0 in model_info should bypass budget.""" @@ -65,9 +64,7 @@ class TestUnmappedModelBudgetEnforcement: ] ) result = _is_model_cost_zero(model="free-model", llm_router=router) - assert ( - result is True - ), "Explicitly free model should bypass budget (return True)" + assert result is True, "Explicitly free model should bypass budget (return True)" def test_known_paid_model_enforces_budget(self): """A model in the cost map with non-zero costs should enforce budget.""" @@ -101,9 +98,7 @@ class TestUnmappedModelBudgetEnforcement: ] ) result = _is_model_cost_zero(model="free-via-params", llm_router=router) - assert ( - result is True - ), "Model with explicit cost=0 in litellm_params should bypass budget" + assert result is True, "Model with explicit cost=0 in litellm_params should bypass budget" def test_cache_invalidates_on_in_place_pricing_update(self): """ @@ -285,9 +280,12 @@ class TestUnmappedModelBudgetEnforcement: "An aliased PTU group must not be read as free" ) - def test_hidden_model_group_alias_enforces_budget(self): - """A hidden alias keeps budget enforced: get_model_group_info() returns None for it, - so the cost is unknown before the configuration gate is reached.""" + def test_hidden_model_group_alias_to_free_model_bypasses_budget(self): + """A hidden alias to an explicitly free group bypasses budget, like the group itself. + + ``get_model_group_info`` returns None for hidden aliases, so the alias must be + resolved to its target group before the cost lookup. + """ router = Router( model_list=[ { @@ -304,7 +302,22 @@ class TestUnmappedModelBudgetEnforcement: model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}}, ) - assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is False + assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is True + + def test_hidden_model_group_alias_to_paid_model_enforces_budget(self): + """A hidden alias to a priced group keeps budget enforced.""" + router = Router( + model_list=[ + { + "model_name": "paid-model", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-fake"}, + "model_info": {"id": "paid-model-id"}, + }, + ], + model_group_alias={"hidden-paid-alias": {"model": "paid-model", "hidden": True}}, + ) + + assert _is_model_cost_zero(model="hidden-paid-alias", llm_router=router) is False def test_dangling_model_group_alias_enforces_budget(self): """An alias pointing at a group that does not exist keeps budget enforced.""" @@ -326,6 +339,115 @@ class TestUnmappedModelBudgetEnforcement: assert _is_model_cost_zero(model="dangling-alias", llm_router=router) is False + def test_repointed_hidden_alias_does_not_reuse_cached_free_result(self): + """Repointing a hidden alias from a free group to a paid group re-evaluates the cost. + + ``Router.update_settings`` is the one runtime path that rewrites the alias map (the + proxy's config update applies ``router_settings`` through it), so the cached verdict + has to drop there. + """ + router = Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "free-model-id"}, + }, + { + "model_name": "paid-model", + "litellm_params": {"model": "gpt-3.5-turbo", "api_key": "sk-fake"}, + "model_info": {"id": "paid-model-id"}, + }, + ], + model_group_alias={"hidden-alias": {"model": "free-model", "hidden": True}}, + ) + + assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is True + router.update_settings(model_group_alias={"hidden-alias": {"model": "paid-model", "hidden": True}}) + assert _is_model_cost_zero(model="hidden-alias", llm_router=router) is False + + @pytest.mark.parametrize("alias_name_first", [True, False]) + def test_alias_shadowing_a_real_group_gives_each_name_its_own_verdict(self, alias_name_first: bool): + """An alias whose name is also a real PTU-priced group never shares a verdict with its target. + + The verdict is cached per requested name, so whichever name is asked first, the free target + stays free and the shadowed PTU name stays enforced. + """ + router = Router( + model_list=[ + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "free-model-id"}, + }, + { + "model_name": "ptu-model", + "litellm_params": { + "model": "azure/ptu-deployment", + "api_base": "https://fake.openai.azure.com", + "api_key": "sk-fake", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "ptu-model-id", "ptu_count": 100, "cost_per_ptu_per_hour": 2.0}, + }, + ], + model_group_alias={"ptu-model": "free-model"}, + ) + order = ("ptu-model", "free-model") if alias_name_first else ("free-model", "ptu-model") + expected = {"ptu-model": False, "free-model": True} + + assert [_is_model_cost_zero(model=name, llm_router=router) for name in order] == [ + expected[name] for name in order + ] + assert [_is_model_cost_zero(model=name, llm_router=router) for name in order] == [ + expected[name] for name in order + ], "the cached verdicts must match the first evaluation" + + def test_alias_chain_through_a_priced_group_enforces_budget(self): + """An alias to a group that is itself an alias key resolves one hop, like the router does. + + The router serves ``chain-smart`` with the real ``chain-legacy`` deployment, which is priced, + so following the second hop to the free group would waive the budget for a paid call. + """ + router = Router( + model_list=[ + { + "model_name": "chain-legacy", + "litellm_params": { + "model": "gpt-3.5-turbo", + "api_key": "sk-fake", + "input_cost_per_token": 0.0000002, + "output_cost_per_token": 0.0000012, + }, + "model_info": {"id": "chain-legacy-id"}, + }, + { + "model_name": "free-model", + "litellm_params": { + "model": "ollama/llama2", + "api_base": "http://localhost:11434", + "input_cost_per_token": 0.0, + "output_cost_per_token": 0.0, + }, + "model_info": {"id": "free-model-id"}, + }, + ], + model_group_alias={"chain-smart": "chain-legacy", "chain-legacy": "free-model"}, + ) + + assert _is_model_cost_zero(model="chain-smart", llm_router=router) is False + def test_handles_router_without_zero_cost_cache_attribute(self): """Tolerate router-like objects (e.g. ``MagicMock`` stand-ins) that do not expose ``_zero_cost_cache`` — the auth check must still diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 823429d7ed6..5d9b37e24fc 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -2278,6 +2278,73 @@ def test_model_group_info_cost_none_for_unpriced_deployment_but_zero_when_declar assert priced.output_cost_per_token is not None and priced.output_cost_per_token > 0 +def _alias_cost_router() -> Router: + return Router( + model_list=[ + { + "model_name": "vllm-free", + "litellm_params": { + "model": "openai/my-vllm-free", + "api_key": "fake", + "api_base": "http://localhost:8000/v1", + "input_cost_per_token": 0, + "output_cost_per_token": 0, + }, + }, + { + "model_name": "gpt-priced", + "litellm_params": {"model": "gpt-4o", "api_key": "fake"}, + }, + ], + model_group_alias={"hidden-free": {"model": "vllm-free", "hidden": True}, "visible": "vllm-free"}, + ) + + +def test_get_model_group_info_include_hidden_resolves_a_hidden_alias(): + router = _alias_cost_router() + + assert router.get_model_group_info(model_group="hidden-free") is None + + hidden: Final = router.get_model_group_info(model_group="hidden-free", include_hidden=True) + assert hidden is not None + assert hidden.model_group == "hidden-free" + assert hidden.input_cost_per_token == 0 + assert hidden.output_cost_per_token == 0 + + +def test_update_settings_model_group_alias_drops_cached_group_info(): + router = _alias_cost_router() + before: Final = router.cached_model_group_info("visible") + assert before is not None and before.input_cost_per_token == 0 + + router.update_settings(model_group_alias={"visible": "gpt-priced"}) + + after: Final = router.cached_model_group_info("visible") + assert after is not None + assert after.input_cost_per_token is not None and after.input_cost_per_token > 0 + + +def test_switch_routing_strategy_installs_lar1_then_restores_the_default_selector(): + router = _alias_cost_router() + + router._switch_routing_strategy( + "lar1", + { + "routing_strategy_args": { + "confidence_threshold_low": 0.1, + "confidence_threshold_medium": 0.3, + "confidence_threshold_high": 0.9, + } + }, + ) + assert router.routing_strategy == "lar1" + assert "async_get_available_deployment" in router.__dict__ + + router._switch_routing_strategy("usage-based-routing-v2", {}) + assert router.lowesttpm_logger_v2 is not None + assert "async_get_available_deployment" not in router.__dict__ + + @pytest.mark.parametrize( "value,expected", [ From 6d73fa6b491089f14a4ebb679abd0a835fa5182a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 16:36:14 -0700 Subject: [PATCH 06/18] fix(vertex_ai): forward system and tools to partner model count_tokens (#43900) * fix(vertex_ai): forward system and tools to partner model count_tokens Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(vertex_ai): avoid mutable token request construction Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(vertex_ai): return partner count_tokens provider errors as values so the proxy falls back locally * test(integration): cover Vertex AI partner count_tokens forwarding and local fallbacks Add wire-level cells for /v1/messages/count_tokens, /utils/token_counter, /v1/responses/input_tokens and the Gemini countTokens route on a Vertex AI Claude deployment: the system prompt and tools reach the partner count-tokens endpoint verbatim, null fields stay out of the body, malformed tools are rejected before any peer call, peer, token-endpoint and connection failures fall back to the local tokenizer unless disable_token_counter is set, generation on the same deployment keeps working, and concurrent bursts survive a peer outage, a slow peer and a worker SIGKILL. The sdk cells cover litellm.acount_tokens the same way. The _support/process.py and _support/client.py harness files are brought to main's content so the self-booting cells read INTEGRATION_PROXY_READY_SECONDS instead of a fixed 70 s boot budget. --------- Co-authored-by: jesus Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- litellm/llms/vertex_ai/common_utils.py | 37 +- .../vertex_ai_partner_models/main.py | 6 + tests/integration/_support/vertex.py | 31 + .../test_vertex_partner_count_tokens_wire.py | 728 ++++++++++++++++++ .../test_vertex_partner_count_tokens_sdk.py | 90 +++ .../vertex_ai/test_vertex_ai_common_utils.py | 143 ++++ 6 files changed, 1027 insertions(+), 8 deletions(-) create mode 100644 tests/integration/_support/vertex.py create mode 100644 tests/integration/providers/test_vertex_partner_count_tokens_wire.py create mode 100644 tests/integration/sdk/test_vertex_partner_count_tokens_sdk.py diff --git a/litellm/llms/vertex_ai/common_utils.py b/litellm/llms/vertex_ai/common_utils.py index 5b5e1403c58..9a0209b87d9 100644 --- a/litellm/llms/vertex_ai/common_utils.py +++ b/litellm/llms/vertex_ai/common_utils.py @@ -1240,6 +1240,9 @@ class VertexAITokenCounter(BaseTokenCounter): ) -> TokenCountResponse | None: import copy + from litellm.llms.vertex_ai.vertex_ai_partner_models.main import ( + VertexAIError as PartnerVertexAIError, + ) from litellm.llms.vertex_ai.vertex_ai_partner_models.main import ( VertexAIPartnerModels, ) @@ -1269,14 +1272,32 @@ class VertexAITokenCounter(BaseTokenCounter): "vertex_ai_credentials" ) - result = await partner_models_handler.count_tokens( - model=model_to_use, - messages=messages or [], - litellm_params=partner_litellm_params, - vertex_project=vertex_project, - vertex_location=vertex_location, - vertex_credentials=vertex_credentials, - ) + try: + result = await partner_models_handler.count_tokens( + model=model_to_use, + messages=messages or [], + litellm_params=partner_litellm_params, + vertex_project=vertex_project, + vertex_location=vertex_location, + vertex_credentials=vertex_credentials, + system=system, + tools=tools, + ) + except (PartnerVertexAIError, httpx.HTTPStatusError) as e: + status_code: Final = e.response.status_code + error_message: Final = e.message if isinstance(e, PartnerVertexAIError) else e.response.text + verbose_logger.warning( + "Vertex AI partner CountTokens API error: status=%s, message=%s", status_code, error_message + ) + return TokenCountResponse( + total_tokens=0, + request_model=request_model, + model_used=model_to_use, + tokenizer_type="vertex_ai_partner_models", + error=True, + error_message=error_message, + status_code=status_code, + ) if result is not None: return TokenCountResponse( diff --git a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py index 40503edbb9e..c855a073648 100644 --- a/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py +++ b/litellm/llms/vertex_ai/vertex_ai_partner_models/main.py @@ -2,6 +2,7 @@ ## API Handler for calling Vertex AI Partner Models from collections.abc import Callable from enum import Enum +from types import MappingProxyType from typing import Final import httpx @@ -263,6 +264,8 @@ class VertexAIPartnerModels(VertexBase): vertex_project=None, vertex_location=None, vertex_credentials=None, + system: object | None = None, + tools: list[dict[str, object]] | None = None, ): """ Count tokens for Vertex AI partner models (Anthropic Claude, Mistral, etc.) @@ -296,6 +299,9 @@ class VertexAIPartnerModels(VertexBase): request_data: Final = { "model": model, "messages": messages, + **MappingProxyType( + {key: value for key, value in (("system", system), ("tools", tools)) if value is not None} + ), } # Prepare litellm_params with credentials diff --git a/tests/integration/_support/vertex.py b/tests/integration/_support/vertex.py new file mode 100644 index 00000000000..af7178fa6e9 --- /dev/null +++ b/tests/integration/_support/vertex.py @@ -0,0 +1,31 @@ +from __future__ import annotations + +import json +from typing import Final + +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa + + +def service_account_json(project: str, token_url: str) -> str: + private_key: Final = ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() + ) + return json.dumps( + { + "type": "service_account", + "project_id": project, + "private_key_id": "scripted", + "private_key": private_key, + "client_email": f"scripted@{project}.iam.gserviceaccount.com", + "client_id": "0", + "auth_uri": f"{token_url}/_oauth/authorize", + "token_uri": f"{token_url}/_oauth/token", + } + ) diff --git a/tests/integration/providers/test_vertex_partner_count_tokens_wire.py b/tests/integration/providers/test_vertex_partner_count_tokens_wire.py new file mode 100644 index 00000000000..5df8dd6e32e --- /dev/null +++ b/tests/integration/providers/test_vertex_partner_count_tokens_wire.py @@ -0,0 +1,728 @@ +import json +import re +import signal +import socket +import threading +import uuid +from collections.abc import Callable, Mapping, Sequence +from concurrent.futures import ThreadPoolExecutor +from contextlib import ExitStack +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, Scenario, eventually +from integration._support.process import owned_proxy_process +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "claude-sonnet-4-6" +_PROJECT: Final = "scripted-project" +_LOCATION: Final = "us-east5" +_MODELS_PATH: Final = f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/anthropic/models" +_COUNT_TARGET: Final = f"{_MODELS_PATH}/count-tokens:rawPredict" +_MESSAGE_TARGET: Final = f"{_MODELS_PATH}/{_BACKEND}:rawPredict" +_STREAM_TARGET: Final = f"{_MODELS_PATH}/{_BACKEND}:streamRawPredict?alt=sse" +_PEER_COUNT: Final = 4242 +_REJECTION: Final = "scripted partner rejection" +_REJECT_TEXT: Final = "The peer must reject this message" +_REPLY_TEXT: Final = "scripted reply" +_OWNED_MODEL: Final = "partner-claude" +_OWNED_UNREACHABLE_MODEL: Final = "partner-claude-unreachable" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") + +_MESSAGES: Final[list[JsonValue]] = [{"role": "user", "content": "Count this message"}] +_SYSTEM: Final = "You are a terse assistant that answers in one sentence" +_SYSTEM_BLOCKS: Final[list[JsonValue]] = [ + {"type": "text", "text": "You are a terse assistant"}, + {"type": "text", "text": "Answer in one sentence"}, +] +_WEATHER_SCHEMA: Final[dict[str, JsonValue]] = { + "type": "object", + "properties": {"city": {"type": "string", "description": "City to look up"}}, + "required": ["city"], +} +_TOOLS: Final[list[JsonValue]] = [ + {"name": "get_weather", "description": "Look up the current weather for a city", "input_schema": _WEATHER_SCHEMA} +] +_OPENAI_TOOLS: Final[list[JsonValue]] = [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Look up the current weather for a city", + "parameters": _WEATHER_SCHEMA, + }, + } +] +_RESPONSES_TOOLS: Final[list[JsonValue]] = [ + { + "type": "function", + "name": "get_weather", + "description": "Look up the current weather for a city", + "parameters": _WEATHER_SCHEMA, + } +] +_PEER_BARE: Final[dict[str, JsonValue]] = {"model": _BACKEND, "messages": _MESSAGES} +_PEER_FULL: Final[dict[str, JsonValue]] = {**_PEER_BARE, "system": _SYSTEM, "tools": _TOOLS} +_GEMINI_BODY: Final[dict[str, JsonValue]] = {"contents": [{"role": "user", "parts": [{"text": "Count this"}]}]} +_GEMINI_MESSAGES: Final[list[JsonValue]] = [{"role": "user", "content": "Count this"}] + +_REPLY: Final[dict[str, JsonValue]] = { + "id": "msg_scripted", + "type": "message", + "role": "assistant", + "model": _BACKEND, + "content": [{"type": "text", "text": _REPLY_TEXT}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 5, "output_tokens": 3}, +} +_EVENTS: Final[tuple[tuple[str, dict[str, JsonValue]], ...]] = ( + ( + "message_start", + { + "type": "message_start", + "message": {**_REPLY, "content": [], "stop_reason": None, "usage": {"input_tokens": 5, "output_tokens": 1}}, + }, + ), + ("content_block_start", {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}), + ( + "content_block_delta", + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": _REPLY_TEXT}}, + ), + ("content_block_stop", {"type": "content_block_stop", "index": 0}), + ( + "message_delta", + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None}, + "usage": {"output_tokens": 3}, + }, + ), + ("message_stop", {"type": "message_stop"}), +) +_SSE: Final = tuple(f"event: {name}\ndata: {json.dumps(data)}\n\n".encode() for name, data in _EVENTS) + + +def _counted(_request: Request) -> Reply: + return Reply(body=json.dumps({"input_tokens": _PEER_COUNT}).encode()) + + +def _rejected(status: int) -> Reply: + return Reply( + status=status, + body=json.dumps({"type": "error", "error": {"type": "invalid_request_error", "message": _REJECTION}}).encode(), + ) + + +def _rejecting(status: int) -> Callable[[Request], Reply]: + def count(_request: Request) -> Reply: + return _rejected(status) + + return count + + +def _anthropic_message(message: JsonValue) -> bool: + return isinstance(message, dict) and message.get("role") in ("user", "assistant") + + +def _anthropic_tool(tool: JsonValue) -> bool: + return isinstance(tool, dict) and isinstance(tool.get("name"), str) and isinstance(tool.get("input_schema"), dict) + + +def _strict(request: Request) -> Reply: + body: Final = _JSON_OBJECT.validate_json(request.body) + messages: Final = body.get("messages") + tools: Final = body.get("tools", []) + accepted: Final = ( + isinstance(messages, list) + and all(map(_anthropic_message, messages)) + and isinstance(body.get("system", ""), (str, list)) + and isinstance(tools, list) + and all(map(_anthropic_tool, tools)) + ) + return _counted(request) if accepted else _rejected(400) + + +def _rejecting_marked_messages(request: Request) -> Reply: + return _rejected(400) if _REJECT_TEXT in request.body.decode() else _counted(request) + + +def _peer(count: Callable[[Request], Reply] = _counted) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + if request.target.endswith("/count-tokens:rawPredict"): + return count(request) + if request.target.endswith(":streamRawPredict?alt=sse"): + return Reply(content_type="text/event-stream", chunks=_SSE) + if request.target.endswith(f"/{_BACKEND}:rawPredict"): + return Reply(body=json.dumps(_REPLY).encode()) + return Reply(status=404, body=json.dumps({"error": f"unscripted target {request.target}"}).encode()) + + return respond + + +def _count_requests(requests: Sequence[Request]) -> tuple[Request, ...]: + return tuple(request for request in requests if "count-tokens" in request.target) + + +def _count_bodies(requests: Sequence[Request], target: str = _COUNT_TARGET) -> tuple[dict[str, JsonValue], ...]: + counts: Final = _count_requests(requests) + for request in counts: + assert (request.method, request.target) == ("POST", target), request.target + assert request.headers["authorization"] == "Bearer scripted-token", request.headers + return tuple(_JSON_OBJECT.validate_json(request.body) for request in counts) + + +def _counted_bodies(wire: Wire) -> tuple[dict[str, JsonValue], ...]: + return _count_bodies(wire.drain()) + + +def _bare(model: str) -> dict[str, JsonValue]: + return {"model": model, "messages": _MESSAGES} + + +def _full(model: str) -> dict[str, JsonValue]: + return {**_bare(model), "system": _SYSTEM, "tools": _TOOLS} + + +def _deployment(gateway: Gateway, scenario: Scenario, api_base: str, **overrides: JsonValue) -> str: + return scenario.model( + **{ + "model": f"vertex_ai/{_BACKEND}", + "api_base": api_base, + "api_key": None, + "vertex_project": _PROJECT, + "vertex_location": _LOCATION, + "vertex_credentials": service_account_json(_PROJECT, gateway.upstream_url.rstrip("/")), + **overrides, + } + ) + + +def _count(gateway: Gateway, body: Mapping[str, JsonValue]) -> httpx.Response: + return gateway.request("POST", "/v1/messages/count_tokens", body) + + +def _payload(response: httpx.Response) -> dict[str, JsonValue]: + assert response.status_code == 200, response.text + return _JSON_OBJECT.validate_json(response.content) + + +def _local_count(gateway: Gateway, body: Mapping[str, JsonValue]) -> int: + response: Final = gateway.request("POST", "/utils/token_counter", body, params={"call_endpoint": "false"}) + payload: Final = _payload(response) + total: Final = payload["total_tokens"] + assert payload["tokenizer_type"] != "vertex_ai_partner_models", response.text + assert isinstance(total, int) and 0 < total != _PEER_COUNT, response.text + return total + + +def _closed_port_url() -> str: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return f"http://127.0.0.1:{reserve.getsockname()[1]}" + + +def _clients(stack: ExitStack, base_url: str, count: int) -> tuple[httpx.Client, ...]: + return tuple( + stack.enter_context(httpx.Client(base_url=base_url, timeout=30, trust_env=False)) for _ in range(count) + ) + + +def _counted_on(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue]: + response: Final = client.post( + "/v1/messages/count_tokens", json=dict(body), headers={"Authorization": f"Bearer {key}"} + ) + return response.status_code, _JSON_OBJECT.validate_json(response.content).get("input_tokens") + + +def _generated_then_counted( + client: httpx.Client, key: str, model: str, body: Mapping[str, JsonValue] +) -> tuple[int, int, JsonValue]: + generated: Final = client.post( + "/v1/messages", + json={ + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": f"Generate before counting {uuid.uuid4().hex}"}], + }, + headers={"Authorization": f"Bearer {key}"}, + ) + return generated.status_code, *_counted_on(client, key, body) + + +def _local_port(client: httpx.Client) -> int: + with client.stream("GET", "/health/liveliness") as response: + port: Final = int(response.extensions["network_stream"].get_extra_info("client_addr")[1]) + response.read() + assert response.status_code == 200, response.text + return port + + +def _counted_or_dropped(client: httpx.Client, key: str, body: Mapping[str, JsonValue]) -> tuple[int, JsonValue] | None: + try: + return _counted_on(client, key, body) + except httpx.TransportError: + return None + + +def _accepted_client_ports(pid: int, proxy_port: int) -> frozenset[int]: + return frozenset( + connection.raddr.port + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.raddr and connection.laddr.port == proxy_port + ) + + +def _owned_config( + path: Path, gateway: Gateway, api_bases: Mapping[str, str], settings: Mapping[str, JsonValue] +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + path.write_text( + yaml.safe_dump( + { + **config, + "model_list": [ + { + "model_name": name, + "litellm_params": { + "model": f"vertex_ai/{_BACKEND}", + "api_base": api_base, + "vertex_project": _PROJECT, + "vertex_location": _LOCATION, + "vertex_credentials": service_account_json(_PROJECT, gateway.upstream_url.rstrip("/")), + }, + } + for name, api_base in api_bases.items() + ], + "litellm_settings": {**config["litellm_settings"], **settings}, + } + ) + ) + return path + + +@pytest.mark.parametrize("system", [_SYSTEM, _SYSTEM_BLOCKS], ids=["string", "blocks"]) +def test_messages_count_tokens_forwards_system_and_tools(gateway: Gateway, system: JsonValue) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, {**_bare(model), "system": system, "tools": _TOOLS}) + assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text + assert _counted_bodies(wire) == ({**_PEER_BARE, "system": system, "tools": _TOOLS},) + + +def test_messages_count_tokens_without_system_or_tools_sends_bare_body(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, _bare(model)) + assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text + assert _counted_bodies(wire) == (_PEER_BARE,) + + +def test_utils_token_counter_call_endpoint_forwards_system_and_tools(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request( + "POST", "/utils/token_counter", _full(model), params={"call_endpoint": "true"} + ) + payload: Final = _payload(response) + assert (payload["total_tokens"], payload["tokenizer_type"]) == (_PEER_COUNT, "vertex_ai_partner_models") + assert (payload["request_model"], payload["model_used"]) == (model, _BACKEND), response.text + assert _counted_bodies(wire) == (_PEER_FULL,) + + +def test_utils_token_counter_local_mode_never_calls_the_peer(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + assert _local_count(gateway, _full(model)) > 0 + assert wire.drain() == () + + +def test_utils_token_counter_falls_back_locally_when_peer_rejects_openai_tools(gateway: Gateway) -> None: + with wire_server(_peer(_strict)) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + body: Final = {**_bare(model), "tools": _OPENAI_TOOLS} + response: Final = gateway.request("POST", "/utils/token_counter", body, params={"call_endpoint": "true"}) + payload: Final = _payload(response) + assert _counted_bodies(wire) == ({**_PEER_BARE, "tools": _OPENAI_TOOLS},) + assert payload["total_tokens"] == _local_count(gateway, body), response.text + assert payload["tokenizer_type"] != "vertex_ai_partner_models", response.text + + +def test_responses_input_tokens_counts_through_the_partner_peer(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request( + "POST", "/v1/responses/input_tokens", {"model": model, "input": "Count this message"} + ) + assert _payload(response) == {"object": "response.input_tokens", "input_tokens": _PEER_COUNT}, response.text + assert _counted_bodies(wire) == (_PEER_BARE,) + + +def test_responses_input_tokens_falls_back_locally_when_peer_rejects(gateway: Gateway) -> None: + with wire_server(_peer(_strict)) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request( + "POST", + "/v1/responses/input_tokens", + {"model": model, "input": "Count this message", "instructions": "Be terse", "tools": _RESPONSES_TOOLS}, + ) + payload: Final = _payload(response) + (sent,) = _counted_bodies(wire) + assert sent["tools"] == _RESPONSES_TOOLS, sent + local: Final = _local_count(gateway, {"model": model, "messages": sent["messages"], "tools": _RESPONSES_TOOLS}) + assert payload == {"object": "response.input_tokens", "input_tokens": local}, response.text + + +def test_gemini_count_tokens_route_reaches_the_partner_peer(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request("POST", f"/v1beta/models/{model}:countTokens", _GEMINI_BODY) + assert "totalTokens" in _payload(response), response.text + assert _counted_bodies(wire) == ({"model": _BACKEND, "messages": _GEMINI_MESSAGES},) + + +def test_gemini_count_tokens_route_falls_back_locally_when_peer_rejects(gateway: Gateway) -> None: + with wire_server(_peer(_rejecting(400))) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request("POST", f"/v1beta/models/{model}:countTokens", _GEMINI_BODY) + payload: Final = _payload(response) + assert _counted_bodies(wire) == ({"model": _BACKEND, "messages": _GEMINI_MESSAGES},) + assert payload["totalTokens"] == _local_count(gateway, {"model": model, "messages": _GEMINI_MESSAGES}) + + +@pytest.mark.parametrize("status", [400, 500, 503]) +def test_messages_count_tokens_falls_back_locally_when_peer_errors(gateway: Gateway, status: int) -> None: + with wire_server(_peer(_rejecting(status))) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, _full(model)) + payload: Final = _payload(response) + assert _counted_bodies(wire) == (_PEER_FULL,) + assert payload == {"input_tokens": _local_count(gateway, _full(model))}, response.text + + +def test_messages_count_tokens_falls_back_locally_when_token_endpoint_rejects(gateway: Gateway) -> None: + def respond(request: Request) -> Reply: + if request.target == "/_oauth/token": + return Reply( + status=400, + body=json.dumps({"error": "invalid_grant", "error_description": "scripted refusal"}).encode(), + ) + return _peer()(request) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = _deployment( + gateway, scenario, wire.url, vertex_credentials=service_account_json(_PROJECT, wire.url) + ) + response: Final = _count(gateway, _full(model)) + payload: Final = _payload(response) + targets: Final = frozenset(request.target for request in wire.drain()) + assert targets == {"/_oauth/token"}, targets + assert payload == {"input_tokens": _local_count(gateway, _full(model))}, response.text + + +def test_messages_count_tokens_falls_back_locally_when_peer_is_unreachable(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, _closed_port_url()) + response: Final = _count(gateway, _full(model)) + assert _payload(response) == {"input_tokens": _local_count(gateway, _full(model))}, response.text + + +@pytest.mark.parametrize( + "tools", + [5, "", "x" * 5120, ["get_weather"]], + ids=["int", "empty_string", "5kb_string", "list_of_strings"], +) +def test_messages_count_tokens_rejects_malformed_tools_without_calling_the_peer( + gateway: Gateway, tools: JsonValue +) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + refused: Final = _count(gateway, {**_bare(model), "tools": tools}) + assert 400 <= refused.status_code < 600, refused.text + assert "input_tokens" not in refused.text, refused.text + assert wire.drain() == () + assert _generated_then_counted(gateway.client, gateway.key, model, _bare(model)) == (200, 200, _PEER_COUNT) + assert _counted_bodies(wire) == (_PEER_BARE,) + + +def test_messages_count_tokens_forwards_an_empty_tools_list(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, {**_bare(model), "tools": []}) + assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text + assert _counted_bodies(wire) == ({**_PEER_BARE, "tools": []},) + + +def test_messages_count_tokens_forwards_an_empty_system_string(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, {**_bare(model), "system": ""}) + assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text + assert _counted_bodies(wire) == ({**_PEER_BARE, "system": ""},) + + +def test_messages_count_tokens_leaves_null_system_and_tools_out(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, {**_bare(model), "system": None, "tools": None}) + assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text + assert _counted_bodies(wire) == (_PEER_BARE,) + + +def test_messages_count_tokens_falls_back_locally_when_peer_rejects_a_non_text_system(gateway: Gateway) -> None: + with wire_server(_peer(_strict)) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, {**_bare(model), "system": 5}) + payload: Final = _payload(response) + assert _counted_bodies(wire) == ({**_PEER_BARE, "system": 5},) + assert payload == {"input_tokens": _local_count(gateway, _bare(model))}, response.text + + +def test_messages_count_tokens_forwards_a_5kb_system_verbatim(gateway: Gateway) -> None: + system: Final = "Answer in one sentence. " * 214 + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, {**_bare(model), "system": system}) + assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text + assert _counted_bodies(wire) == ({**_PEER_BARE, "system": system},) + + +def test_messages_count_tokens_duplicate_system_and_tools_keys_forward_one_value_each(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + fields: Final = f'"system": {json.dumps(_SYSTEM)}, "tools": {json.dumps(_TOOLS)}' + response: Final = gateway.client.post( + "/v1/messages/count_tokens", + content=f'{{"model": "{model}", "messages": {json.dumps(_MESSAGES)}, {fields}, {fields}}}', + headers={"Authorization": f"Bearer {gateway.key}", "Content-Type": "application/json"}, + ) + assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text + (sent,) = _count_requests(wire.drain()) + assert _count_bodies((sent,)) == (_PEER_FULL,) + assert (sent.body.count(b'"system"'), sent.body.count(b'"tools"')) == (1, 1), sent.body + + +def test_messages_count_tokens_unauthenticated_request_never_reaches_the_peer(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request( + "POST", "/v1/messages/count_tokens", _full(model), key="sk-not-a-key-this-proxy-issued" + ) + assert response.status_code == 401, response.text + assert "input_tokens" not in response.text, response.text + assert wire.drain() == () + + +@pytest.mark.parametrize("fields", [{}, {"messages": []}], ids=["missing", "empty"]) +def test_messages_count_tokens_without_messages_is_rejected(gateway: Gateway, fields: dict[str, JsonValue]) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = _count(gateway, {"model": model, "system": _SYSTEM, "tools": _TOOLS, **fields}) + assert response.status_code == 400, response.text + assert "messages parameter is required" in response.text, response.text + assert wire.drain() == () + + +def test_count_tokens_location_override_targets_the_count_region(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment( + gateway, scenario, wire.url, vertex_location="global", vertex_count_tokens_location="europe-west1" + ) + response: Final = _count(gateway, _bare(model)) + assert _payload(response) == {"input_tokens": _PEER_COUNT}, response.text + target: Final = _COUNT_TARGET.replace(f"/locations/{_LOCATION}/", "/locations/europe-west1/") + assert _count_bodies(wire.drain(), target) == (_PEER_BARE,) + + +def test_messages_count_tokens_repeated_request_reaches_the_peer_each_time(gateway: Gateway) -> None: + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + answers: Final = tuple(_payload(_count(gateway, _bare(model))) for _ in range(2)) + assert answers == ({"input_tokens": _PEER_COUNT},) * 2 + assert _counted_bodies(wire) == (_PEER_BARE,) * 2 + + +@pytest.mark.timeout(240) # boots an owned two-worker proxy with litellm_settings.disable_token_counter +def test_disabled_token_counter_surfaces_provider_failures_instead_of_counting_locally( + gateway: Gateway, tmp_path: Path +) -> None: + with wire_server(_peer(_rejecting_marked_messages)) as wire: + config: Final = _owned_config( + tmp_path / "disabled-token-counter.yaml", + gateway, + {_OWNED_MODEL: wire.url, _OWNED_UNREACHABLE_MODEL: _closed_port_url()}, + {"disable_token_counter": True}, + ) + with owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2) as owned: + counted: Final = _count(owned.gateway, _full(_OWNED_MODEL)) + assert _payload(counted) == {"input_tokens": _PEER_COUNT}, counted.text + rejected: Final = _count( + owned.gateway, {**_full(_OWNED_MODEL), "messages": [{"role": "user", "content": _REJECT_TEXT}]} + ) + assert rejected.status_code == 400, rejected.text + assert _REJECTION in rejected.text and "input_tokens" not in rejected.text, rejected.text + unreachable: Final = _count(owned.gateway, _full(_OWNED_UNREACHABLE_MODEL)) + assert 500 <= unreachable.status_code < 600, unreachable.text + assert "input_tokens" not in unreachable.text, unreachable.text + local: Final = owned.gateway.request( + "POST", "/utils/token_counter", _full(_OWNED_MODEL), params={"call_endpoint": "false"} + ) + assert local.status_code == 503, local.text + assert len(_counted_bodies(wire)) == 2 + + +def test_peer_outage_between_concurrent_waves_falls_back_then_recovers(gateway: Gateway) -> None: + with ExitStack() as stack: + clients: Final = _clients(stack, str(gateway.client.base_url), 8) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + scenario: Final = stack.enter_context(gateway.scenario()) + with wire_server(_peer()) as wire: + port: Final = int(wire.url.rsplit(":", 1)[1]) + model: Final = _deployment(gateway, scenario, wire.url) + body: Final = _full(model) + local: Final = _local_count(gateway, body) + + def generate_then_count(client: httpx.Client) -> tuple[int, int, JsonValue]: + return _generated_then_counted(client, gateway.key, model, body) + + assert tuple(pool.map(generate_then_count, clients)) == ((200, 200, _PEER_COUNT),) * len(clients) + assert _counted_bodies(wire) == (_PEER_FULL,) * len(clients) + outage: Final = tuple(pool.map(lambda client: _counted_on(client, gateway.key, body), clients)) + assert outage == ((200, local),) * len(clients) + with wire_server(_peer(), port=port) as revived: + assert tuple(pool.map(generate_then_count, clients)) == ((200, 200, _PEER_COUNT),) * len(clients) + assert _counted_bodies(revived) == (_PEER_FULL,) * len(clients) + + +def test_slow_peer_holds_concurrent_counts_without_stalling_the_proxy(gateway: Gateway) -> None: + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + + def hold(request: Request) -> Reply: + held.put(request.target) + assert release.wait(timeout=20), "Held count was never released" + return _counted(request) + + with ExitStack() as stack: + clients: Final = _clients(stack, str(gateway.client.base_url), 6) + wire: Final = stack.enter_context(wire_server(_peer(hold))) + scenario: Final = stack.enter_context(gateway.scenario()) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + model: Final = _deployment(gateway, scenario, wire.url) + try: + futures: Final = tuple( + pool.submit(_generated_then_counted, client, gateway.key, model, _full(model)) for client in clients + ) + eventually(held.qsize, lambda size: size == len(clients), seconds=30) + assert gateway.request("GET", "/health/liveliness").status_code == 200 + assert _local_count(gateway, _full(model)) > 0 + assert not any(future.done() for future in futures) + finally: + release.set() + assert tuple(future.result(timeout=30) for future in futures) == ((200, 200, _PEER_COUNT),) * len(clients) + assert len(_counted_bodies(wire)) == len(clients) + + +@pytest.mark.timeout(300) # boots an owned two-worker proxy, kills one worker, and waits for its replacement +def test_worker_sigkill_mid_burst_leaves_the_sibling_counting(gateway: Gateway, tmp_path: Path) -> None: + held: Final[SimpleQueue[str]] = SimpleQueue() + release: Final = threading.Event() + + def hold(request: Request) -> Reply: + held.put(request.target) + assert release.wait(timeout=60), "Held count was never released" + return _counted(request) + + with ExitStack() as stack: + wire: Final = stack.enter_context(wire_server(_peer(hold))) + stack.callback(release.set) + config: Final = _owned_config(tmp_path / "worker-kill.yaml", gateway, {_OWNED_MODEL: wire.url}, {}) + owned: Final = stack.enter_context(owned_proxy_process(gateway, tmp_path, {}, config=config, workers=2)) + proxy_url: Final = owned.gateway.client.base_url + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + clients: Final = _clients(stack, str(proxy_url), 12) + pool: Final = stack.enter_context(ThreadPoolExecutor(max_workers=len(clients))) + stack.callback(release.set) + ports: Final = tuple(_local_port(client) for client in clients) + futures: Final = tuple( + pool.submit(_counted_or_dropped, client, gateway.key, _full(_OWNED_MODEL)) for client in clients + ) + eventually(held.qsize, lambda size: size == len(clients), seconds=30) + shares: Final = {pid: _accepted_client_ports(pid, proxy_url.port or 0) & frozenset(ports) for pid in workers} + assert sum(map(len, shares.values())) == len(clients), shares + victim: Final = min((pid for pid in workers if shares[pid]), key=lambda pid: len(shares[pid])) + psutil.Process(victim).send_signal(signal.SIGKILL) + release.set() + results: Final = tuple(future.result(timeout=60) for future in futures) + for port, result in zip(ports, results, strict=True): + assert result == (None if port in shares[victim] else (200, _PEER_COUNT)), (port, result, shares) + second_wave: Final = _clients(stack, str(proxy_url), 6) + assert tuple(_counted_on(client, gateway.key, _full(_OWNED_MODEL)) for client in second_wave) == ( + (200, _PEER_COUNT), + ) * len(second_wave) + assert len(_counted_bodies(wire)) == len(clients) + len(second_wave) + eventually(lambda: len(_STARTED_WORKER.findall(owned.log.read_text())), lambda started: started >= 3, 120) + assert owned.process.poll() is None + + +@pytest.mark.parametrize("stream", [False, True], ids=["non_stream", "stream"]) +def test_chat_completions_on_the_same_deployment_still_generate(gateway: Gateway, stream: bool) -> None: + marker: Final = f"chat control {uuid.uuid4().hex}" + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": marker}], + "stream": stream, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + assert _REPLY_TEXT in response.text, response.text + assert not stream or response.text.rstrip().endswith("data: [DONE]"), response.text + (sent,) = wire.drain() + assert sent.target == (_STREAM_TARGET if stream else _MESSAGE_TARGET), sent.target + assert marker in sent.body.decode(), sent.body + + +def test_messages_endpoint_on_the_same_deployment_still_generates(gateway: Gateway) -> None: + marker: Final = f"messages control {uuid.uuid4().hex}" + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": marker}]}, + ) + assert response.status_code == 200, response.text + assert _REPLY_TEXT in response.text, response.text + (sent,) = wire.drain() + assert sent.target == _MESSAGE_TARGET, sent.target + assert marker in sent.body.decode(), sent.body + + +def test_responses_endpoint_on_the_same_deployment_still_generates(gateway: Gateway) -> None: + marker: Final = f"responses control {uuid.uuid4().hex}" + with wire_server(_peer()) as wire, gateway.scenario() as scenario: + model: Final = _deployment(gateway, scenario, wire.url) + response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": marker}) + assert response.status_code == 200, response.text + assert _REPLY_TEXT in response.text, response.text + (sent,) = wire.drain() + assert sent.target == _MESSAGE_TARGET, sent.target + assert marker in sent.body.decode(), sent.body diff --git a/tests/integration/sdk/test_vertex_partner_count_tokens_sdk.py b/tests/integration/sdk/test_vertex_partner_count_tokens_sdk.py new file mode 100644 index 00000000000..e59d3f9264a --- /dev/null +++ b/tests/integration/sdk/test_vertex_partner_count_tokens_sdk.py @@ -0,0 +1,90 @@ +import json +from collections.abc import Callable +from typing import Final + +import litellm +import pytest +from integration._support.vertex import service_account_json +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "claude-sonnet-4-6" +_PROJECT: Final = "scripted-project" +_LOCATION: Final = "us-east5" +_COUNT_TARGET: Final = ( + f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/anthropic/models/count-tokens:rawPredict" +) +_PEER_COUNT: Final = 4242 +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_MESSAGES: Final[list[dict[str, str]]] = [{"role": "user", "content": "Count this message"}] +_SYSTEM: Final = "You are a terse assistant that answers in one sentence" +_TOOLS: Final[list[dict[str, JsonValue]]] = [ + { + "name": "get_weather", + "description": "Look up the current weather for a city", + "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, + } +] + + +def _peer(status: int) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + if request.target == "/_oauth/token": + token: Final = {"access_token": "scripted-token", "token_type": "Bearer", "expires_in": 3600} + return Reply(body=json.dumps(token).encode()) + if status == 200: + return Reply(body=json.dumps({"input_tokens": _PEER_COUNT}).encode()) + rejection: Final = {"type": "error", "error": {"type": "invalid_request_error", "message": "scripted"}} + return Reply(status=status, body=json.dumps(rejection).encode()) + + return respond + + +def _count_requests(requests: tuple[Request, ...]) -> tuple[Request, ...]: + return tuple(request for request in requests if "count-tokens" in request.target) + + +@pytest.fixture +def vertex_environment(monkeypatch: pytest.MonkeyPatch) -> Callable[[str], None]: + def configure(token_url: str) -> None: + monkeypatch.setenv("VERTEXAI_PROJECT", _PROJECT) + monkeypatch.setenv("VERTEXAI_LOCATION", _LOCATION) + monkeypatch.setenv("VERTEXAI_CREDENTIALS", service_account_json(_PROJECT, token_url)) + + return configure + + +async def test_acount_tokens_forwards_system_and_tools_to_the_partner_peer( + vertex_environment: Callable[[str], None], +) -> None: + with wire_server(_peer(200)) as wire: + vertex_environment(wire.url) + counted: Final = await litellm.acount_tokens( + model=f"vertex_ai/{_BACKEND}", messages=_MESSAGES, tools=_TOOLS, system=_SYSTEM, api_base=wire.url + ) + (sent,) = _count_requests(wire.drain()) + assert (sent.method, sent.target, sent.headers["authorization"]) == ( + "POST", + _COUNT_TARGET, + "Bearer scripted-token", + ) + assert _JSON_OBJECT.validate_json(sent.body) == { + "model": _BACKEND, + "messages": _MESSAGES, + "system": _SYSTEM, + "tools": _TOOLS, + } + assert (counted.total_tokens, counted.tokenizer_type) == (_PEER_COUNT, "vertex_ai_partner_models"), counted + + +async def test_acount_tokens_falls_back_to_the_local_tokenizer_when_the_peer_rejects( + vertex_environment: Callable[[str], None], +) -> None: + with wire_server(_peer(400)) as wire: + vertex_environment(wire.url) + counted: Final = await litellm.acount_tokens( + model=f"vertex_ai/{_BACKEND}", messages=_MESSAGES, tools=_TOOLS, system=_SYSTEM, api_base=wire.url + ) + assert len(_count_requests(wire.drain())) == 1 + assert counted.tokenizer_type == "local_tokenizer", counted + assert counted.total_tokens > 0 and counted.total_tokens != _PEER_COUNT, counted diff --git a/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py b/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py index 04a7ee451c4..86eb26a4c15 100644 --- a/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py +++ b/tests/unit/llms/vertex_ai/test_vertex_ai_common_utils.py @@ -1229,6 +1229,149 @@ async def test_vertex_ai_token_counter_routes_partner_models(): assert result.tokenizer_type == "vertex_ai_partner_models" +@pytest.mark.asyncio +async def test_vertex_ai_token_counter_forwards_system_and_tools_to_partner_request(): + from typing import Final + from unittest.mock import AsyncMock, patch + + from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter + from litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens import handler + + class FakeResponse: + status_code = 200 + + def json(self) -> dict[str, int]: + return {"input_tokens": 37} + + class FakeHttpClient: + posted_bodies: tuple[dict[str, object], ...] = () + + async def post( + self, + url: str, + headers: dict[str, str], + json: dict[str, object], + timeout: float, + ) -> FakeResponse: + self.posted_bodies = (*self.posted_bodies, json) + return FakeResponse() + + fake_http_client: Final = FakeHttpClient() + counter: Final = VertexAITokenCounter() + model: Final = "claude-opus-5-5" + messages: Final = [{"role": "user", "content": "Hello"}] + system: Final = "Follow the system instructions" + tools: Final = [ + { + "name": "lookup", + "description": "Look up a value", + "input_schema": {"type": "object", "properties": {}}, + } + ] + deployment: Final = { + "litellm_params": { + "vertex_project": "test-project", + "vertex_location": "us-east5", + } + } + + with ( + patch.object(handler, "get_async_httpx_client", return_value=fake_http_client), + patch.object( + handler.VertexAIPartnerModelsTokenCounter, + "_ensure_access_token_async", + new=AsyncMock(return_value=("fake-token", "test-project")), + ), + ): + with_optional_fields: Final = await counter.count_tokens( + model_to_use=model, + messages=messages, + contents=None, + deployment=deployment, + system=system, + tools=tools, + ) + without_optional_fields: Final = await counter.count_tokens( + model_to_use=model, + messages=messages, + contents=None, + deployment=deployment, + ) + + assert fake_http_client.posted_bodies == ( + {"model": model, "messages": messages, "system": system, "tools": tools}, + {"model": model, "messages": messages}, + ) + assert with_optional_fields is not None + assert with_optional_fields.total_tokens == 37 + assert without_optional_fields is not None + assert without_optional_fields.total_tokens == 37 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("provider_failure", "expected_status", "expected_message"), + [ + ("http_400", 400, 'tools.0: Input tag "function" does not match any of the expected tags'), + ("credentials", 500, "could not resolve credentials"), + ], +) +async def test_vertex_ai_token_counter_returns_partner_provider_error_as_value( + provider_failure: str, expected_status: int, expected_message: str +): + from typing import Final + from unittest.mock import AsyncMock, patch + + import httpx + + from litellm.llms.custom_httpx.http_handler import MaskedHTTPStatusError + from litellm.llms.vertex_ai.common_utils import VertexAITokenCounter + from litellm.llms.vertex_ai.vertex_ai_partner_models.count_tokens import handler + + class RejectingHttpClient: + async def post( + self, + url: str, + headers: dict[str, str], + json: dict[str, object], + timeout: float, + ) -> None: + request: Final = httpx.Request("POST", url) + response: Final = httpx.Response(400, text=expected_message, request=request) + raise MaskedHTTPStatusError( + httpx.HTTPStatusError("Client error '400 Bad Request'", request=request, response=response), + message=expected_message, + text=expected_message, + ) + + access_token: Final = ( + AsyncMock(side_effect=ValueError(expected_message)) + if provider_failure == "credentials" + else AsyncMock(return_value=("fake-token", "test-project")) + ) + with ( + patch.object(handler, "get_async_httpx_client", return_value=RejectingHttpClient()), + patch.object(handler.VertexAIPartnerModelsTokenCounter, "_ensure_access_token_async", new=access_token), + ): + result: Final = await VertexAITokenCounter().count_tokens( + model_to_use="claude-opus-5-5", + messages=[{"role": "user", "content": "Hello"}], + contents=None, + deployment={"litellm_params": {"vertex_project": "test-project", "vertex_location": "us-east5"}}, + request_model="vertex-claude", + tools=[{"type": "function", "function": {"name": "lookup", "parameters": {}}}], + ) + + assert result is not None + assert result.error is True + assert result.status_code == expected_status + assert result.error_message is not None + assert expected_message in result.error_message + assert result.total_tokens == 0 + assert result.request_model == "vertex-claude" + assert result.tokenizer_type == "vertex_ai_partner_models" + + @pytest.mark.asyncio async def test_vertex_ai_token_counter_uses_count_tokens_location(): """ From 21ecd0af557597d7fc2cad48a3aae2e8c2f1094b Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 23:43:10 +0000 Subject: [PATCH 07/18] fix(traces): reject conflicting spend aliases and unrelated HTTP siblings (#44456) Co-authored-by: Yujong Lee --- .../crates/traces-clickhouse/src/span_row.rs | 2 +- .../tests/queries/support.rs | 3 +- .../normalize/instrumentation/http_client.rs | 41 ++++--- .../crates/traces/src/normalize/mod.rs | 3 + litellm-rust/crates/traces/src/otlp/span.rs | 4 +- .../crates/traces/src/resolve/resolution.rs | 16 ++- .../crates/traces/src/resolve/spend.rs | 50 ++++++++- litellm-rust/crates/traces/tests/captures.rs | 12 +- .../traces/tests/normalization_formats.rs | 3 +- litellm-rust/crates/traces/tests/normalize.rs | 3 + litellm-rust/crates/traces/tests/resolve.rs | 106 ++++++++++++++++-- 11 files changed, 204 insertions(+), 39 deletions(-) diff --git a/litellm-rust/crates/traces-clickhouse/src/span_row.rs b/litellm-rust/crates/traces-clickhouse/src/span_row.rs index c3b652129a4..0189dbff2c8 100644 --- a/litellm-rust/crates/traces-clickhouse/src/span_row.rs +++ b/litellm-rust/crates/traces-clickhouse/src/span_row.rs @@ -211,7 +211,7 @@ fn request_id(evidence: &CallEvidence) -> &str { .flatten() .find_map(|key| match key { CallKey::ProviderResponse(id) => Some(id.as_str()), - CallKey::LiteLlmRequest(_) | CallKey::Transport => None, + CallKey::LiteLlmRequest(_) | CallKey::Transport | CallKey::GatewayAttempt => None, }) .unwrap_or_default() } diff --git a/litellm-rust/crates/traces-clickhouse/tests/queries/support.rs b/litellm-rust/crates/traces-clickhouse/tests/queries/support.rs index a9f42119fb4..371424c63d6 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/queries/support.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/queries/support.rs @@ -158,7 +158,8 @@ fn span_row(span: &DecodedSpan, team: &str, key: &str) -> BTreeMap Some(id.as_str()), - litellm_traces::CallKey::Transport => None, + litellm_traces::CallKey::Transport + | litellm_traces::CallKey::GatewayAttempt => None, }) .unwrap_or_default() ), diff --git a/litellm-rust/crates/traces/src/normalize/instrumentation/http_client.rs b/litellm-rust/crates/traces/src/normalize/instrumentation/http_client.rs index f7b563535c1..1ee6c24b9d4 100644 --- a/litellm-rust/crates/traces/src/normalize/instrumentation/http_client.rs +++ b/litellm-rust/crates/traces/src/normalize/instrumentation/http_client.rs @@ -12,23 +12,30 @@ const SCOPES: [&str; 7] = [ ]; pub(super) fn matches(context: &SpanContext<'_>) -> bool { - SCOPES.contains(&context.scope) - || (context.scope == "litellm.gateway.client" - && context.name == "gateway.request" - && context - .attributes - .get("litellm.gateway.attempt") - .is_some_and(|value| value == "true") - && context - .attributes - .get("http.request.method") - .is_some_and(|value| value == "POST")) + SCOPES.contains(&context.scope) || matches_gateway_attempt(context) } -pub(super) fn adjust(facts: SpanFacts) -> SpanFacts { +fn matches_gateway_attempt(context: &SpanContext<'_>) -> bool { + context.scope == "litellm.gateway.client" + && context.name == "gateway.request" + && context + .attributes + .get("litellm.gateway.attempt") + .is_some_and(|value| value == "true") + && context + .attributes + .get("http.request.method") + .is_some_and(|value| value == "POST") +} + +pub(super) fn adjust(context: &SpanContext<'_>, facts: SpanFacts) -> SpanFacts { SpanFacts { role: Some(RoleEvidence::Declared(ObservationType::Framework)), - calls: CallEvidence::complete(CallKey::Transport), + calls: CallEvidence::complete(if matches_gateway_attempt(context) { + CallKey::GatewayAttempt + } else { + CallKey::Transport + }), ..facts } } @@ -42,7 +49,11 @@ impl Rule for HttpClient { fn integration(&self, _: &SpanContext<'_>) -> Option { None } - fn adjust(&self, _: &SpanContext<'_>, extraction: super::Extraction) -> super::Extraction { - extraction.map_facts(adjust) + fn adjust( + &self, + context: &SpanContext<'_>, + extraction: super::Extraction, + ) -> super::Extraction { + extraction.map_facts(|facts| adjust(context, facts)) } } diff --git a/litellm-rust/crates/traces/src/normalize/mod.rs b/litellm-rust/crates/traces/src/normalize/mod.rs index 7d17f53daa8..c3ad76fde02 100644 --- a/litellm-rust/crates/traces/src/normalize/mod.rs +++ b/litellm-rust/crates/traces/src/normalize/mod.rs @@ -53,6 +53,7 @@ pub enum CallKey { ProviderResponse(String), /// The span is the HTTP request itself; LiteLLM logs its `traceparent` span id. Transport, + GatewayAttempt, } impl fmt::Display for CallKey { @@ -61,6 +62,7 @@ impl fmt::Display for CallKey { Self::LiteLlmRequest(id) => write!(formatter, "litellm_request:{id}"), Self::ProviderResponse(id) => write!(formatter, "provider_response:{id}"), Self::Transport => formatter.write_str("transport:"), + Self::GatewayAttempt => formatter.write_str("gateway_attempt:"), } } } @@ -77,6 +79,7 @@ impl FromStr for CallKey { Ok(Self::LiteLlmRequest(id.to_owned())) } Some(("transport", "")) => Ok(Self::Transport), + Some(("gateway_attempt", "")) => Ok(Self::GatewayAttempt), _ => Err(crate::InvalidCallKey), } } diff --git a/litellm-rust/crates/traces/src/otlp/span.rs b/litellm-rust/crates/traces/src/otlp/span.rs index 58aba3c68b9..0261e4aeb07 100644 --- a/litellm-rust/crates/traces/src/otlp/span.rs +++ b/litellm-rust/crates/traces/src/otlp/span.rs @@ -171,7 +171,9 @@ fn decoded_span( crate::CallKey::LiteLlmRequest(id) | crate::CallKey::ProviderResponse(id) => { id.len() + size_of::() } - crate::CallKey::Transport => size_of::(), + crate::CallKey::Transport | crate::CallKey::GatewayAttempt => { + size_of::() + } }) .sum::() + normalized.model.as_ref().map_or(0, String::len) diff --git a/litellm-rust/crates/traces/src/resolve/resolution.rs b/litellm-rust/crates/traces/src/resolve/resolution.rs index 1b8ffffa812..b469fb95019 100644 --- a/litellm-rust/crates/traces/src/resolve/resolution.rs +++ b/litellm-rust/crates/traces/src/resolve/resolution.rs @@ -134,11 +134,13 @@ impl<'a> Resolution<'a> { ) } - /// The request attempts a model call made: its transport descendants, or, for bridges that - /// emit the request beside the call instead of under it, transport siblings inside the call's - /// time window when the call is the only model call under that parent. fn transports(&self, call: usize) -> Vec { - let is_transport = |index: &usize| self.row(*index).call_keys.contains(&CallKey::Transport); + let is_transport = |index: &usize| { + self.row(*index) + .call_keys + .iter() + .any(|key| matches!(key, CallKey::Transport | CallKey::GatewayAttempt)) + }; let nested: Vec = self .graph .descendants(call) @@ -162,7 +164,11 @@ impl<'a> Resolution<'a> { let call_end_ns = call_start_ns + i128::from(call_row.duration_ns); siblings .into_iter() - .filter(is_transport) + .filter(|sibling| { + self.row(*sibling) + .call_keys + .contains(&CallKey::GatewayAttempt) + }) .filter(|sibling| { let transport = self.row(*sibling); let transport_start_ns = i128::from(transport.start_ns); diff --git a/litellm-rust/crates/traces/src/resolve/spend.rs b/litellm-rust/crates/traces/src/resolve/spend.rs index 34c90a7d12c..ce23362b7f8 100644 --- a/litellm-rust/crates/traces/src/resolve/spend.rs +++ b/litellm-rust/crates/traces/src/resolve/spend.rs @@ -48,7 +48,9 @@ impl SpendLookup { trace_ids: sorted( keys() .filter_map(|(row, key)| match key { - CallKey::Transport if !row.trace_id.is_empty() => { + CallKey::Transport | CallKey::GatewayAttempt + if !row.trace_id.is_empty() => + { Some(row.trace_id.clone()) } _ => None, @@ -80,6 +82,21 @@ impl Ownership<'_> { pub(super) type Requests<'a> = Vec<&'a SpendRow>; +#[derive(Clone, Copy, Eq, Ord, PartialEq, PartialOrd)] +enum KeyFamily { + GatewayCall, + ProviderResponse, + Transport, +} + +fn key_family(key: &CallKey) -> KeyFamily { + match key { + CallKey::LiteLlmRequest(_) => KeyFamily::GatewayCall, + CallKey::ProviderResponse(_) => KeyFamily::ProviderResponse, + CallKey::Transport | CallKey::GatewayAttempt => KeyFamily::Transport, + } +} + pub(super) enum KeyMatch<'a> { Missing, Unique(&'a SpendRow), @@ -174,7 +191,7 @@ fn matches<'a>( && (spend.litellm_call_id == *id || (spend.litellm_call_id.is_empty() && spend.request_id == *id)) } - CallKey::Transport => { + CallKey::Transport | CallKey::GatewayAttempt => { !row.trace_id.is_empty() && !row.span_id.is_empty() && spend.trace_id == row.trace_id @@ -216,12 +233,37 @@ pub(super) fn requests<'a>( && anchored .iter() .all(|request| request.litellm_call_id.is_empty()); - let matches = keyed + let aliases: Vec<_> = keyed .into_iter() .filter(|(key, requests)| { !(legacy_rows && requests.is_empty() && matches!(key, CallKey::LiteLlmRequest(_))) }) - .map(|(_, requests)| KeyMatch::new(requests)) + .collect(); + let families: BTreeSet<_> = aliases.iter().map(|(key, _)| key_family(key)).collect(); + let compatible_rows: Vec> = families + .into_iter() + .map(|family| { + aliases + .iter() + .filter(|(key, _)| key_family(key) == family) + .flat_map(|(_, requests)| requests.iter().map(|request| request.identity())) + .collect() + }) + .collect(); + let matches = aliases + .into_iter() + .map(|(_, requests)| { + KeyMatch::new( + requests + .into_iter() + .filter(|request| { + compatible_rows + .iter() + .all(|family| family.contains(&request.identity())) + }) + .collect(), + ) + }) .collect(); match evidence.kind() { CallEvidenceKind::Complete => SpendEvidence::Complete(matches), diff --git a/litellm-rust/crates/traces/tests/captures.rs b/litellm-rust/crates/traces/tests/captures.rs index 30f84086a9b..fc3527f6253 100644 --- a/litellm-rust/crates/traces/tests/captures.rs +++ b/litellm-rust/crates/traces/tests/captures.rs @@ -172,7 +172,7 @@ fn trace_span(span: DecodedSpan) -> TraceSpansRow { .flatten() .find_map(|key| match key { CallKey::ProviderResponse(id) => Some(id.clone()), - CallKey::LiteLlmRequest(_) | CallKey::Transport => None, + CallKey::LiteLlmRequest(_) | CallKey::Transport | CallKey::GatewayAttempt => None, }) .unwrap_or_default(); TraceSpansRow { @@ -359,7 +359,11 @@ fn unrelated_sibling_transport_leaves_cost_unchanged( let calls: Vec<_> = rows .iter() .filter(|row| { - row.kind == ObservationType::Llm && !row.call_keys.contains(&CallKey::Transport) + row.kind == ObservationType::Llm + && !row + .call_keys + .iter() + .any(|key| matches!(key, CallKey::Transport | CallKey::GatewayAttempt)) }) .cloned() .collect(); @@ -408,7 +412,9 @@ fn redundant_genai_response_id_keeps_call_evidence( .iter() .filter_map(|key| match key { CallKey::ProviderResponse(id) => Some(id.clone()), - CallKey::LiteLlmRequest(_) | CallKey::Transport => None, + CallKey::LiteLlmRequest(_) | CallKey::Transport | CallKey::GatewayAttempt => { + None + } }) .collect(); (!response_ids.is_empty()).then(|| { diff --git a/litellm-rust/crates/traces/tests/normalization_formats.rs b/litellm-rust/crates/traces/tests/normalization_formats.rs index dfa5753978d..0cb2152ab19 100644 --- a/litellm-rust/crates/traces/tests/normalization_formats.rs +++ b/litellm-rust/crates/traces/tests/normalization_formats.rs @@ -746,6 +746,7 @@ fn transport_contract_keeps_independent_call_ids(span: Span) { #[case::unrelated_scope("custom", "gateway.request", "true", "POST", false)] #[case::unrelated_span("litellm.gateway.client", "step", "true", "POST", false)] #[case::missing_contract("litellm.gateway.client", "gateway.request", "", "POST", false)] +#[case::disabled_contract("litellm.gateway.client", "gateway.request", "false", "POST", false)] #[case::unrelated_method("litellm.gateway.client", "gateway.request", "true", "GET", false)] fn gateway_attempt_contract_requires_recorded_request_boundary( span: Span, @@ -774,7 +775,7 @@ fn gateway_attempt_contract_requires_recorded_request_boundary( decoded.normalized.calls, if complete { CallEvidence::Complete(std::collections::BTreeSet::from([ - CallKey::Transport, + CallKey::GatewayAttempt, gateway, ])) } else { diff --git a/litellm-rust/crates/traces/tests/normalize.rs b/litellm-rust/crates/traces/tests/normalize.rs index cab1bb8d08f..244b814cdd6 100644 --- a/litellm-rust/crates/traces/tests/normalize.rs +++ b/litellm-rust/crates/traces/tests/normalize.rs @@ -255,6 +255,7 @@ fn llamaindex_wrapped_responses_keep_provider_call_keys(#[case] body: &[u8]) { #[case::request(litellm_traces::CallKey::LiteLlmRequest("request:with:colons".to_owned()))] #[case::response(litellm_traces::CallKey::ProviderResponse("response:with:colons".to_owned()))] #[case::transport(litellm_traces::CallKey::Transport)] +#[case::gateway_attempt(litellm_traces::CallKey::GatewayAttempt)] fn call_keys_round_trip_through_storage(#[case] key: litellm_traces::CallKey) { assert_eq!( key.to_string().parse::().unwrap(), @@ -272,6 +273,8 @@ fn call_keys_round_trip_through_storage(#[case] key: litellm_traces::CallKey) { #[case::missing_response("provider_response:")] #[case::missing_request("litellm_request:")] #[case::transport_id("transport:unexpected")] +#[case::gateway_attempt_separator("gateway_attempt")] +#[case::gateway_attempt_id("gateway_attempt:unexpected")] #[case::unknown("unknown:id")] fn malformed_call_keys_are_rejected_at_the_boundary(#[case] encoded: &str) { assert!(encoded.parse::().is_err()); diff --git a/litellm-rust/crates/traces/tests/resolve.rs b/litellm-rust/crates/traces/tests/resolve.rs index 009c9b8f929..6e987cd4970 100644 --- a/litellm-rust/crates/traces/tests/resolve.rs +++ b/litellm-rust/crates/traces/tests/resolve.rs @@ -579,7 +579,7 @@ fn sibling_transports_belong_to_the_only_model_call_under_their_parent( 10, ); transport.trace_id = "trace".into(); - transport.call_keys = vec!["transport:".parse().unwrap()]; + transport.call_keys = vec![litellm_traces::CallKey::GatewayAttempt]; transport.call_evidence = Some(litellm_traces::CallEvidenceKind::Complete); let mut rows = vec![ owned( @@ -605,11 +605,15 @@ fn sibling_transports_belong_to_the_only_model_call_under_their_parent( } #[rstest] -#[case::without_tool_http_sibling(None, Some(0.5))] -#[case::after_call(Some((200, 10)), Some(0.5))] -#[case::inside_call_without_spend(Some((10, 10)), None)] +#[case::without_tool_http_sibling(None, false, litellm_traces::CallKey::Transport, Some(0.5))] +#[case::after_call(Some((200, 10)), false, litellm_traces::CallKey::Transport, Some(0.5))] +#[case::inside_call_without_spend(Some((10, 10)), false, litellm_traces::CallKey::Transport, Some(0.5))] +#[case::inside_call_with_unrelated_spend(Some((10, 10)), true, litellm_traces::CallKey::Transport, Some(0.5))] +#[case::missing_gateway_attempt(Some((10, 10)), false, litellm_traces::CallKey::GatewayAttempt, None)] fn sibling_transport_does_not_lose_model_call_spend( #[case] transport_timing: Option<(i64, u64)>, + #[case] unrelated_spend: bool, + #[case] key: litellm_traces::CallKey, #[case] expected: Option, ) { let call = owned( @@ -637,24 +641,110 @@ fn sibling_transport_does_not_lose_model_call_spend( ]; let rows: Vec<_> = base_rows .into_iter() - .chain(transport_timing.into_iter().map(|(start, duration)| { + .chain(transport_timing.map(|(start, duration)| { let mut transport = at( row("tool-http", "step", "GET", "framework", ""), start, duration, ); transport.trace_id = "trace".into(); - transport.call_keys = vec![litellm_traces::CallKey::Transport]; + transport.call_keys = vec![key]; transport.call_evidence = Some(litellm_traces::CallEvidenceKind::Complete); owned(transport, "team", "", "key") })) .collect(); - let logged = spend("chatcmpl-1", "chatcmpl-1", "team", "", "key", 0.5); - let trace = resolve_trace("trace", "ref", &rows, &[logged]).unwrap(); + let logs: Vec<_> = std::iter::once(spend("chatcmpl-1", "chatcmpl-1", "team", "", "key", 0.5)) + .chain(unrelated_spend.then(|| SpendByResponseIdsRow { + trace_id: "trace".into(), + span_id: "tool-http".into(), + ..spend("unrelated", "unrelated", "team", "", "key", 0.75) + })) + .collect(); + let trace = resolve_trace("trace", "ref", &rows, &logs).unwrap(); assert_eq!(trace.summary.spend, expected); assert_eq!(trace.agents[0].spend, expected); } +#[rstest] +#[case::agreeing_ids( + litellm_traces::CallKey::Transport, + "call-a", + Some("response-a"), + Some(0.25) +)] +#[case::conflicting_gateway_id(litellm_traces::CallKey::Transport, "call-b", None, None)] +#[case::conflicting_response_id( + litellm_traces::CallKey::Transport, + "call-a", + Some("response-b"), + None +)] +#[case::conflicting_gateway_and_response( + litellm_traces::CallKey::Transport, + "call-b", + Some("response-b"), + None +)] +#[case::agreeing_gateway_attempt( + litellm_traces::CallKey::GatewayAttempt, + "call-a", + Some("response-a"), + Some(0.25) +)] +#[case::conflicting_gateway_attempt(litellm_traces::CallKey::GatewayAttempt, "call-b", None, None)] +fn gateway_attempt_identifiers_must_match_one_spend_row( + #[case] transport: litellm_traces::CallKey, + #[case] call_id: &str, + #[case] response_id: Option<&str>, + #[case] expected: Option, +) { + let keys = [ + transport, + litellm_traces::CallKey::LiteLlmRequest(call_id.into()), + ] + .into_iter() + .chain(response_id.map(|id| litellm_traces::CallKey::ProviderResponse(id.into()))) + .collect(); + let rows = [ + owned( + row("agent", "", "agent", "agent", "agent"), + "team", + "", + "key", + ), + owned(llm("call", "agent", "agent", ""), "team", "", "key"), + owned( + TraceSpansRow { + trace_id: "trace".into(), + call_keys: keys, + call_evidence: Some(litellm_traces::CallEvidenceKind::Complete), + ..row("attempt", "call", "gateway.request", "framework", "") + }, + "team", + "", + "key", + ), + ]; + let logs = [ + SpendByResponseIdsRow { + litellm_call_id: "call-a".into(), + trace_id: "trace".into(), + span_id: "attempt".into(), + ..spend("request-a", "response-a", "team", "", "key", 0.25) + }, + SpendByResponseIdsRow { + litellm_call_id: "call-b".into(), + trace_id: "trace".into(), + span_id: "other-attempt".into(), + ..spend("request-b", "response-b", "team", "", "key", 0.5) + }, + ]; + let trace = resolve_trace("trace", "ref", &rows, &logs).unwrap(); + assert_eq!(trace.summary.spend, expected); + assert_eq!(trace.agents[0].spend, expected); + assert_eq!(trace.spans[2].spend, expected); +} + #[rstest] #[case::legacy_row("", Some(0.5))] #[case::other_call("other-call", None)] From 4d30f8c59be965891cd3d6335855b06b12ad7e20 Mon Sep 17 00:00:00 2001 From: "berriai-litellm-provider-info-sync[bot]" <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 16:44:13 -0700 Subject: [PATCH 08/18] chore(cost-map): add azure_ai/kimi-k2-thinking retirement date from the Azure retired models page (#44455) Price-Sync: litellm-providers Co-authored-by: berriai-litellm-provider-info-sync[bot] <328147090+berriai-litellm-provider-info-sync[bot]@users.noreply.github.com> --- litellm/model_prices_and_context_window_backup.json | 1 + model_prices_and_context_window.json | 1 + 2 files changed, 2 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 27495443a57..da62a2c2677 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -79983,6 +79983,7 @@ "supports_web_search": false }, "azure_ai/kimi-k2-thinking": { + "deprecation_date": "2026-03-29", "input_cost_per_token": 6e-07, "litellm_provider": "azure_ai", "max_input_tokens": 262144, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 27495443a57..da62a2c2677 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -79983,6 +79983,7 @@ "supports_web_search": false }, "azure_ai/kimi-k2-thinking": { + "deprecation_date": "2026-03-29", "input_cost_per_token": 6e-07, "litellm_provider": "azure_ai", "max_input_tokens": 262144, From f0eda6d2a6f83af8c592a25b83e0021fcc823f37 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 17:05:47 -0700 Subject: [PATCH 09/18] fix(health): probe Bedrock Mantle Claude deployments over the Anthropic Messages API (#44419) * fix(health): probe Bedrock Mantle Claude deployments over the Anthropic Messages API Bedrock Mantle serves Claude ids only on /anthropic/v1/messages, but health checks probed every chat-mode deployment over /v1/chat/completions, so a bedrock_mantle Claude deployment showed unhealthy while real /v1/messages traffic to it succeeded Add an anthropic_messages health check mode and make it the default for bedrock_mantle Claude models. An explicit model_info.mode still wins, and /health/test_connection and the Add Model form accept the new mode * fix(health): resolve the test connection mode from the deployment when the request omits it The Admin UI model page sent the mode /model/info had filled in from the cost map back as the probe mode, so Test Connection on a Bedrock Mantle Claude deployment still went over chat completions. The page now forwards only the row's id, and /health/test_connection resolves a missing mode the way /health does: the stored model_info.mode, then the mode the provider requires, then the cost map. * fix(health): resolve an omitted ahealth_check mode the way the proxy does * fix(health): test connection honors a stored mode only for the stored model and rejects a non-string mode A request that selects a stored deployment and sends a different litellm_params.model now resolves the probe mode from that model instead of the stored model_info.mode. A litellm_params.mode that is not a string answers 400 instead of 500. The Bedrock Mantle rule that Claude models are probed over the Messages API moves into the provider package. * fix(health): shape test connection probe params for the model the request probes A request that selects a stored deployment by id and overrides the model resolved its probe mode from the overridden model but still injected max_tokens from the stored mode, so an embedding override of an anthropic_messages deployment failed with a Mistral 422 extra_forbidden * fix(health): report an early ahealth_check failure as itself, not as a missing mode With the mode resolved automatically when the caller omits it, a failure before that resolution (no model, a non-string model, a provider that does not resolve) was wrapped as "Missing mode", a hint that pointed at the wrong fix and dropped raw_request_typed_dict from the result. Every failure now returns the same shape. --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- .../health_check_helpers.py | 34 ++++ .../litellm_core_utils/health_check_utils.py | 1 + litellm/llms/bedrock_mantle/common_utils.py | 6 +- litellm/main.py | 26 +-- litellm/proxy/health_check.py | 33 +++- .../health_endpoints/_health_endpoints.py | 44 ++++- .../test_health_check_helpers.py | 161 ++++++++++++++++ .../health_endpoints/test_health_endpoints.py | 179 +++++++++++++++++- .../proxy/test_health_check_max_tokens.py | 117 +++++++++++- .../components/add_model/add_model_modes.tsx | 1 + .../src/components/model_info_view.test.tsx | 29 +++ .../src/components/model_info_view.tsx | 2 - .../src/components/networking.tsx | 2 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 4 +- 14 files changed, 596 insertions(+), 43 deletions(-) diff --git a/litellm/litellm_core_utils/health_check_helpers.py b/litellm/litellm_core_utils/health_check_helpers.py index 41965404351..fba59b0b983 100644 --- a/litellm/litellm_core_utils/health_check_helpers.py +++ b/litellm/litellm_core_utils/health_check_helpers.py @@ -24,6 +24,32 @@ IMAGE_EDIT_HEALTH_CHECK_PROMPT: Final = ( "Add a small yellow star in the top right corner of this simple drawing of a blue circle on a white background" ) +ANTHROPIC_MESSAGES_HEALTH_CHECK_MAX_TOKENS: Final = 16 + + +def native_health_check_mode(model: str, custom_llm_provider: str | None) -> Literal["anthropic_messages"] | None: + if custom_llm_provider != "bedrock_mantle": + return None + from litellm.llms.bedrock_mantle.common_utils import mantle_health_check_mode + + return mantle_health_check_mode(model) + + +def _cost_map_mode(model: str) -> str | None: + import litellm + from litellm.litellm_core_utils.health_check_utils import OPTIONAL_STR + + return OPTIONAL_STR.validate_python(litellm.model_cost.get(model, {}).get("mode")) + + +def default_health_check_mode(requested_model: str, model: str, custom_llm_provider: str) -> str: + return ( + native_health_check_mode(model=model, custom_llm_provider=custom_llm_provider) + or _cost_map_mode(requested_model) + or _cost_map_mode(model) + or "chat" + ) + def get_image_file_for_health_check() -> bytes: """Return the image used for health checks.""" @@ -167,6 +193,7 @@ class HealthCheckHelpers: "realtime", "batch", "responses", + "anthropic_messages", "ocr", "evaluation", ], @@ -254,6 +281,13 @@ class HealthCheckHelpers: **_filter_model_params(model_params=model_params), input=prompt or "test", ), + "anthropic_messages": lambda: litellm.anthropic_messages( + **{ + "max_tokens": ANTHROPIC_MESSAGES_HEALTH_CHECK_MAX_TOKENS, + "messages": [{"role": "user", "content": prompt or "test"}], + **model_params, + } + ), "ocr": lambda: litellm.aocr( **_filter_model_params(model_params=model_params), document=_ocr_health_check_document(model=model, custom_llm_provider=custom_llm_provider), diff --git a/litellm/litellm_core_utils/health_check_utils.py b/litellm/litellm_core_utils/health_check_utils.py index 7fe2d830f1e..5205e9f1284 100644 --- a/litellm/litellm_core_utils/health_check_utils.py +++ b/litellm/litellm_core_utils/health_check_utils.py @@ -9,6 +9,7 @@ from pydantic import TypeAdapter from litellm.types.decisions import DecisionsCallParams DECISIONS_CALL_PARAMS: Final[TypeAdapter[DecisionsCallParams]] = TypeAdapter(DecisionsCallParams) +OPTIONAL_STR: Final[TypeAdapter[str | None]] = TypeAdapter(str | None) def _filter_model_params(model_params: dict) -> dict: diff --git a/litellm/llms/bedrock_mantle/common_utils.py b/litellm/llms/bedrock_mantle/common_utils.py index 5a43da95604..232d6dcf70f 100644 --- a/litellm/llms/bedrock_mantle/common_utils.py +++ b/litellm/llms/bedrock_mantle/common_utils.py @@ -14,7 +14,7 @@ global state. import re from collections.abc import Mapping -from typing import Final +from typing import Final, Literal from botocore.exceptions import ( CredentialRetrievalError, @@ -131,6 +131,10 @@ def is_mantle_claude_model(model: str) -> bool: return "claude" in model.lower() +def mantle_health_check_mode(model: str) -> Literal["anthropic_messages"] | None: + return "anthropic_messages" if is_mantle_claude_model(model) else None + + def mantle_supports_responses(model: str | None, model_cost: dict) -> bool: """Whether a Bedrock Mantle model can serve the native Responses API. diff --git a/litellm/main.py b/litellm/main.py index b24b35166ed..6cfc2b8af55 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -8725,7 +8725,7 @@ def speech( async def ahealth_check( model_params: dict, - mode: str | None = "chat", + mode: str | None = None, prompt: str | None = None, input: list | None = None, ): @@ -8740,7 +8740,8 @@ async def ahealth_check( } """ from litellm.litellm_core_utils.cached_imports import get_litellm_logging_class - from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers + from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers, default_health_check_mode + from litellm.litellm_core_utils.health_check_utils import OPTIONAL_STR # Use cached import helper to lazy-load Logging class (only loads when function is called) Logging: Final = get_litellm_logging_class() @@ -8765,28 +8766,25 @@ async def ahealth_check( ) ######################################################### try: - model: str | None = model_params.get("model", None) - if model is None: + requested_model: Final = OPTIONAL_STR.validate_python(model_params.get("model", None)) + if requested_model is None: raise Exception("model not set") - if model in litellm.model_cost and mode is None: - mode = litellm.model_cost[model].get("mode") - custom_llm_provider_from_params: Final = model_params.get("custom_llm_provider", None) api_base_from_params: Final = model_params.get("api_base", None) api_key_from_params: Final = model_params.get("api_key", None) model, custom_llm_provider, _, _ = get_llm_provider( - model=model, + model=requested_model, custom_llm_provider=custom_llm_provider_from_params, api_base=api_base_from_params, api_key=api_key_from_params, ) - if model in litellm.model_cost and mode is None: - mode = litellm.model_cost[model].get("mode") model_params["cache"] = {"no-cache": True} # don't used cached responses for making health check calls - mode = mode or "chat" + mode = mode or default_health_check_mode( + requested_model=requested_model, model=model, custom_llm_provider=custom_llm_provider + ) if "*" in model: return await HealthCheckHelpers.ahealth_check_wildcard_models( model=model, @@ -8815,12 +8813,6 @@ async def ahealth_check( if isinstance(stack_trace, str): stack_trace = stack_trace[:1000] - if mode is None: - return { - "error": f"error:{e}. Missing `mode`. Set the `mode` for the model - https://docs.litellm.ai/docs/proxy/health#embedding-models \nstacktrace: {stack_trace}", - "exception": e, - } - error_to_return: Final = str(e) + "\nstack trace: " + stack_trace raw_request_typed_dict: Final = litellm_logging_obj.model_call_details.get("raw_request_typed_dict") diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index 88c65e954fa..d5d0123cb6a 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -26,6 +26,7 @@ from litellm.constants import ( DEFAULT_HEALTH_CHECK_PROMPT, HEALTH_CHECK_TIMEOUT_SECONDS, ) +from litellm.litellm_core_utils.health_check_helpers import native_health_check_mode from litellm.router_utils.auto_router_model_naming import ( StrategyRouterDependency, classify_strategy_router_model, @@ -69,18 +70,29 @@ HEALTH_DISPLAY_PARAMS: Final = ( # endpoints that reject unknown fields with 400 "Unknown parameter: # 'max_tokens'". Allow-list so new modes are safe by default. # Per-deployment override: `model_info.health_check_supports_max_tokens`. -_MAX_TOKEN_SUPPORT_MODES: Final[frozenset[str]] = frozenset({"chat", "completion", "responses"}) +_MAX_TOKEN_SUPPORT_MODES: Final[frozenset[str]] = frozenset({"chat", "completion", "responses", "anthropic_messages"}) -def _resolve_health_check_mode(model_info: Mapping[str, object], litellm_params: Mapping[str, object]) -> str | None: +def _native_health_check_mode(model: str, provider_param: object) -> str | None: + try: + resolved_model, custom_llm_provider, _, _ = litellm.get_llm_provider( + model=model, custom_llm_provider=provider_param if isinstance(provider_param, str) else None + ) + except Exception: + return None + return native_health_check_mode(model=resolved_model, custom_llm_provider=custom_llm_provider) + + +def resolve_health_check_mode(model_info: Mapping[str, object], litellm_params: Mapping[str, object]) -> str | None: """ Effective mode for a deployment's health-check probe. - Prefers operator-set `model_info.mode`; otherwise resolves it from the model - cost map, which understands `bedrock/` and cross-region inference-profile - prefixes (`us.`, `eu.`, `apac.`). Without this, non-chat Bedrock deployments - (e.g. embeddings) are probed as chat, so `max_tokens` is injected and the - request 400s on "extraneous key [max_tokens]". + Prefers operator-set `model_info.mode`; then the mode the provider requires for + that model family (Bedrock Mantle serves Claude ids on the Messages API only); + otherwise resolves it from the model cost map, which understands `bedrock/` and + cross-region inference-profile prefixes (`us.`, `eu.`, `apac.`). Without this, + non-chat Bedrock deployments (e.g. embeddings) are probed as chat, so + `max_tokens` is injected and the request 400s on "extraneous key [max_tokens]". """ explicit_mode: Final = model_info.get("mode") if isinstance(explicit_mode, str): @@ -88,6 +100,9 @@ def _resolve_health_check_mode(model_info: Mapping[str, object], litellm_params: model: Final = litellm_params.get("model") if not isinstance(model, str): return None + native_mode: Final = _native_health_check_mode(model, litellm_params.get("custom_llm_provider")) + if native_mode is not None: + return native_mode try: return litellm.get_model_info(model=model).get("mode") except Exception: @@ -518,7 +533,7 @@ async def _run_model_health_check(model: dict): if _is_strategy_router_deployment(litellm_params): return {} - mode: Final = _resolve_health_check_mode( + mode: Final = resolve_health_check_mode( model_info, litellm_params, # any-ok: untyped router config dict ) @@ -768,7 +783,7 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di - updates the `voice` param with the `health_check_voice` for `audio_speech` mode if it exists Doc: https://docs.litellm.ai/docs/proxy/health#text-to-speech-models - for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID, and pins `custom_llm_provider` to `bedrock` (only when the deployment hasn't already set one, so an explicit `bedrock_converse` survives) so the bare model id still resolves to the provider (e.g. cross-region ids like `us.cohere.embed-v4:0`) """ - mode: Final = _resolve_health_check_mode( + mode: Final = resolve_health_check_mode( model_info, litellm_params, # any-ok: untyped router config dict ) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 0f389518f7b..8804a190d4d 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -12,6 +12,7 @@ from typing import Any, Final, Literal, TypedDict, cast import fastapi from fastapi import APIRouter, Depends, HTTPException, Request, Response, status +from pydantic import TypeAdapter from typing_extensions import ReadOnly import litellm @@ -58,6 +59,7 @@ from litellm.proxy.health_check import ( deployments_targeted_by_name, health_check_filter_kwargs_from_general_settings, perform_health_check, + resolve_health_check_mode, run_with_timeout, ) from litellm.proxy.middleware.admission_control_middleware import ( @@ -173,6 +175,24 @@ def _config_base_for_health_check( return {key: value for key, value in config_params.items() if key not in _CONFIG_CONNECTION_FIELDS} +def _model_info_for_mode_resolution( + model_info: Mapping[str, object], stored_params: Mapping[str, object], request_params: Mapping[str, object] +) -> Mapping[str, object]: + stored_model: Final = stored_params.get("model") + if stored_model is None or request_params.get("model") in (None, stored_model): + return model_info + return {key: value for key, value in model_info.items() if key != "mode"} + + +def _string_mode_or_bad_request(params_mode: object) -> str | None: + if params_mode is None or isinstance(params_mode, str): + return params_mode + raise HTTPException( + status_code=status.HTTP_400_BAD_REQUEST, + detail={"error": f"litellm_params.mode must be a string, got {type(params_mode).__name__}"}, + ) + + def get_callback_identifier(callback): """ Get the callback identifier string, handling both strings and objects. @@ -203,6 +223,7 @@ def get_callback_identifier(callback): router: Final = APIRouter() +_OBJECT_MAPPING: Final = TypeAdapter(Mapping[str, object]) services = ( Literal[ "slack_budget_alerts", @@ -2033,11 +2054,16 @@ async def test_model_connection( "rerank", "realtime", "responses", + "anthropic_messages", "ocr", ] | None = fastapi.Body( None, - description="The mode to test the model with. If not provided, auto-detected from model capabilities.", + description=( + "The mode to test the model with. If not provided, resolved the way /health does: the deployment's " + "model_info.mode (only while the request tests the deployment's own model), then the mode the " + "provider requires for that model, then the model cost map." + ), ), litellm_params: dict = fastapi.Body( None, @@ -2188,8 +2214,13 @@ async def test_model_connection( } resolved_model_info: Final = loaded_model_info if loaded_model_info is not None else model_info + probe_model_info: Final = _model_info_for_mode_resolution( + _OBJECT_MAPPING.validate_python(resolved_model_info or {}), + stored_params=_OBJECT_MAPPING.validate_python(config_litellm_params), + request_params=_OBJECT_MAPPING.validate_python(request_litellm_params), + ) litellm_params = _update_litellm_params_for_health_check( - model_info=resolved_model_info or {}, + model_info=dict(probe_model_info), litellm_params=litellm_params, ) @@ -2204,12 +2235,17 @@ async def test_model_connection( prisma_client=prisma_client, premium_user=premium_user, ) - mode = mode or litellm_params.pop("mode", None) + raw_params_mode: Final[object] = litellm_params.pop("mode", None) + probe_mode: Final = ( + mode + or _string_mode_or_bad_request(raw_params_mode) + or resolve_health_check_mode(probe_model_info, _OBJECT_MAPPING.validate_python(litellm_params)) + ) result: Final = await run_with_timeout( litellm.ahealth_check( model_params=litellm_params, - mode=mode, + mode=probe_mode, prompt="test from litellm", input=["test from litellm"], ), diff --git a/tests/unit/litellm_core_utils/test_health_check_helpers.py b/tests/unit/litellm_core_utils/test_health_check_helpers.py index 941e44feb26..e8a81d8ac87 100644 --- a/tests/unit/litellm_core_utils/test_health_check_helpers.py +++ b/tests/unit/litellm_core_utils/test_health_check_helpers.py @@ -16,6 +16,8 @@ from litellm.constants import LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME from litellm.litellm_core_utils.health_check_helpers import ( IMAGE_EDIT_HEALTH_CHECK_PROMPT, HealthCheckHelpers, + default_health_check_mode, + native_health_check_mode, ) from litellm.main import ahealth_check from litellm.proxy._types import UserAPIKeyAuth @@ -647,3 +649,162 @@ async def test_ahealth_check_probes_strands_through_decisions_without_mode( assert "error" not in result, result assert upstream.called assert "authorization" not in upstream.calls[0].request.headers + + +@pytest.mark.parametrize( + ("model", "custom_llm_provider", "expected"), + ( + ("anthropic.claude-haiku-4-5", "bedrock_mantle", "anthropic_messages"), + ("Anthropic.Claude-Opus-5-5", "bedrock_mantle", "anthropic_messages"), + ("openai.gpt-oss-120b", "bedrock_mantle", None), + ("us.anthropic.claude-haiku-4-5-20251001-v1:0", "bedrock", None), + ("claude-haiku-4-5", "anthropic", None), + ("anthropic.claude-haiku-4-5", None, None), + ), +) +def test_native_health_check_mode_is_messages_only_for_mantle_claude( + model: str, custom_llm_provider: str | None, expected: str | None +) -> None: + assert native_health_check_mode(model=model, custom_llm_provider=custom_llm_provider) == expected + + +def test_default_health_check_mode_prefers_the_native_surface_over_the_cost_map( + monkeypatch: pytest.MonkeyPatch, +) -> None: + monkeypatch.setattr(litellm, "model_cost", {"anthropic.claude-haiku-4-5": {"mode": "chat"}}) + + assert ( + default_health_check_mode( + requested_model="bedrock_mantle/anthropic.claude-haiku-4-5", + model="anthropic.claude-haiku-4-5", + custom_llm_provider="bedrock_mantle", + ) + == "anthropic_messages" + ) + + +@pytest.mark.parametrize( + ("model_cost", "expected"), + ( + ({"bedrock_mantle/openai.gpt-oss-120b": {"mode": "responses"}}, "responses"), + ({"openai.gpt-oss-120b": {"mode": "completion"}}, "completion"), + ( + { + "bedrock_mantle/openai.gpt-oss-120b": {"mode": "responses"}, + "openai.gpt-oss-120b": {"mode": "completion"}, + }, + "responses", + ), + ({}, "chat"), + ), +) +def test_default_health_check_mode_falls_back_to_cost_map_then_chat( + model_cost: dict[str, dict[str, str]], expected: str, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setattr(litellm, "model_cost", model_cost) + + assert ( + default_health_check_mode( + requested_model="bedrock_mantle/openai.gpt-oss-120b", + model="openai.gpt-oss-120b", + custom_llm_provider="bedrock_mantle", + ) + == expected + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode_kwargs", [{}, {"mode": None}], ids=["omitted", "explicit_none"]) +async def test_ahealth_check_probes_mantle_claude_through_messages_without_mode( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, + mode_kwargs: dict[str, None], +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("https://bedrock-mantle.us-east-2.api.aws/anthropic/v1/messages").respond( + json={ + "id": "msg_health", + "type": "message", + "role": "assistant", + "model": "anthropic.claude-haiku-4-5", + "content": [{"type": "text", "text": "pong"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 3, "output_tokens": 1}, + } + ) + + result: Final = await ahealth_check( + { + "model": "bedrock_mantle/anthropic.claude-haiku-4-5", + "api_key": "test-bearer", + "aws_region_name": "us-east-2", + }, + prompt="test from litellm", + **mode_kwargs, + ) + + assert "error" not in result, result + assert upstream.call_count == 1 + sent: Final = json.loads(upstream.calls.last.request.content) + assert sent["model"] == "anthropic.claude-haiku-4-5" + assert sent["max_tokens"] == 16 + assert sent["messages"] == [{"role": "user", "content": "test from litellm"}] + + +@pytest.mark.asyncio +async def test_ahealth_check_anthropic_messages_mode_keeps_caller_supplied_messages( + monkeypatch: pytest.MonkeyPatch, + respx_mock: respx.MockRouter, +) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + upstream: Final = respx_mock.post("https://bedrock-mantle.us-east-2.api.aws/anthropic/v1/messages").respond( + json={ + "id": "msg_health", + "type": "message", + "role": "assistant", + "model": "anthropic.claude-haiku-4-5", + "content": [{"type": "text", "text": "pong"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 3, "output_tokens": 1}, + } + ) + + result: Final = await ahealth_check( + { + "model": "bedrock_mantle/anthropic.claude-haiku-4-5", + "api_key": "test-bearer", + "aws_region_name": "us-east-2", + "messages": [{"role": "user", "content": "operator probe"}], + "max_tokens": 4, + }, + mode="anthropic_messages", + prompt="test from litellm", + ) + + assert "error" not in result, result + sent: Final = json.loads(upstream.calls.last.request.content) + assert sent["max_tokens"] == 4 + assert sent["messages"] == [{"role": "user", "content": "operator probe"}] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("model_params", "expected_error"), + ( + ({"model": "not-a-provider/some-model"}, "LLM Provider NOT provided"), + ({"api_key": "test-bearer"}, "model not set"), + ), + ids=["unknown_provider", "model_missing"], +) +async def test_ahealth_check_without_mode_reports_the_real_failure( + model_params: dict[str, str], expected_error: str +) -> None: + result: Final = await ahealth_check(model_params, prompt="test from litellm") + + assert expected_error in result["error"], result["error"] + assert "Missing `mode`" not in result["error"] + assert "raw_request_typed_dict" in result diff --git a/tests/unit/proxy/health_endpoints/test_health_endpoints.py b/tests/unit/proxy/health_endpoints/test_health_endpoints.py index f40c33b1e91..513f4a67554 100644 --- a/tests/unit/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/unit/proxy/health_endpoints/test_health_endpoints.py @@ -5,14 +5,14 @@ import time from collections.abc import Iterator, Mapping, Sequence from contextlib import contextmanager from datetime import datetime, timedelta -from types import SimpleNamespace +from types import MappingProxyType, SimpleNamespace from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest import respx -from fastapi import FastAPI +from fastapi import FastAPI, HTTPException from fastapi.testclient import TestClient from prisma.errors import ClientNotConnectedError, HTTPClientClosedError, PrismaError @@ -694,6 +694,181 @@ async def test_test_model_connection_falls_back_to_deployments_zero_without_id() assert model_params.get("api_key") == "fake-key-A" +@contextmanager +def _test_connection_probe( + deployment: Mapping[str, object], +) -> Iterator[AsyncMock]: + from litellm.types.router import Deployment, LiteLLM_Params + + router: Final = MagicMock() + router.get_deployment.side_effect = lambda model_id: ( + Deployment( + model_name=str(deployment["model_name"]), + litellm_params=LiteLLM_Params(**deployment["litellm_params"]), # pyright: ignore[reportArgumentType] # test fixture dict + model_info=deployment["model_info"], # pyright: ignore[reportArgumentType] # test fixture dict + ) + if model_id == deployment["model_info"]["id"] # pyright: ignore[reportIndexIssue] # test fixture dict + else None + ) + ahealth_check: Final = AsyncMock(return_value={"status": "healthy"}) + with ( + patch("litellm.proxy.proxy_server.prisma_client", MagicMock()), + patch("litellm.proxy.proxy_server.llm_router", router), + patch("litellm.proxy.proxy_server.premium_user", False), + patch( + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + AsyncMock(), + ), + patch("litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", ahealth_check), + patch( + "litellm.proxy.health_endpoints._health_endpoints.run_with_timeout", + AsyncMock(return_value={"status": "healthy"}), + ), + ): + yield ahealth_check + + +MANTLE_CLAUDE_DEPLOYMENT: Final = MappingProxyType( + { + "model_name": "claude-haiku-4-5", + "litellm_params": { + "model": "bedrock_mantle/anthropic.claude-haiku-4-5", + "api_key": "fake-mantle-key", + "aws_region_name": "us-east-2", + }, + "model_info": {"id": "mantle-claude-id"}, + } +) + + +@pytest.mark.asyncio +async def test_test_model_connection_without_mode_probes_mantle_claude_over_messages(): + """ + The Admin UI model page sends the row's id and no mode. The probe must then resolve + the mode the way /health does, so a Bedrock Mantle Claude deployment is checked over + the Anthropic Messages API instead of chat completions, which Mantle rejects. + """ + with _test_connection_probe(MANTLE_CLAUDE_DEPLOYMENT) as ahealth_check: + result: Final = await health_test_model_connection( + request=MagicMock(), + mode=None, + litellm_params={"model": "bedrock_mantle/anthropic.claude-haiku-4-5"}, + model_info={"id": "mantle-claude-id"}, + user_api_key_dict=UserAPIKeyAuth(user_id="test-user", token="test-token"), + ) + + assert result["status"] == "success" + assert ahealth_check.call_args.kwargs["mode"] == "anthropic_messages" + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("request_params", "expected_mode"), + [ + ({"model": "bedrock_mantle/anthropic.claude-haiku-4-5"}, "chat"), + ({}, "chat"), + ({"model": "bedrock_mantle/anthropic.claude-sonnet-4-5"}, "anthropic_messages"), + ], + ids=["stored_model", "no_model", "overridden_model"], +) +async def test_test_model_connection_stored_operator_mode_follows_the_stored_model( + request_params: Mapping[str, str], expected_mode: str +): + """ + A mode the operator stored on the deployment is the probe's mode when the request + carries none, ahead of the provider-native rule, but only while the request probes + the deployment's own model. A request that selects the deployment by id and swaps in + another model resolves the mode from that model instead. + """ + deployment: Final = MappingProxyType( + {**MANTLE_CLAUDE_DEPLOYMENT, "model_info": {"id": "mantle-claude-id", "mode": "chat"}} + ) + with _test_connection_probe(deployment) as ahealth_check: + await health_test_model_connection( + request=MagicMock(), + mode=None, + litellm_params=dict(request_params), + model_info={"id": "mantle-claude-id"}, + user_api_key_dict=UserAPIKeyAuth(user_id="test-user", token="test-token"), + ) + + assert ahealth_check.call_args.kwargs["mode"] == expected_mode + + +@pytest.mark.asyncio +async def test_test_model_connection_overridden_model_probe_params_follow_the_probed_model(): + """ + When the request selects a deployment by id and swaps in another model, the probe's + params are shaped for that model, so the stored mode must not inject `max_tokens` + into what is now an embedding probe (Mistral rejects it with a 422 extra_forbidden). + """ + deployment: Final = MappingProxyType( + { + "model_name": "anthropic-claude-haiku-4-5", + "litellm_params": {"model": "anthropic/claude-haiku-4-5", "api_key": "fake-anthropic-key"}, + "model_info": {"id": "anthropic-messages-id", "mode": "anthropic_messages"}, + } + ) + with _test_connection_probe(deployment) as ahealth_check: + await health_test_model_connection( + request=MagicMock(), + mode=None, + litellm_params={"model": "mistral/mistral-embed", "api_key": "fake-mistral-key"}, + model_info={"id": "anthropic-messages-id"}, + user_api_key_dict=UserAPIKeyAuth(user_id="test-user", token="test-token"), + ) + + assert ahealth_check.call_args.kwargs["mode"] == "embedding" + assert "max_tokens" not in ahealth_check.call_args.kwargs["model_params"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("params_mode", [123, ["chat"], {"mode": "chat"}, False], ids=["int", "list", "dict", "bool"]) +async def test_test_model_connection_non_string_params_mode_is_a_bad_request(params_mode: object): + with _test_connection_probe(MANTLE_CLAUDE_DEPLOYMENT) as ahealth_check: + with pytest.raises(HTTPException) as exc_info: + await health_test_model_connection( + request=MagicMock(), + mode=None, + litellm_params={"model": "bedrock_mantle/anthropic.claude-haiku-4-5", "mode": params_mode}, + model_info={"id": "mantle-claude-id"}, + user_api_key_dict=UserAPIKeyAuth(user_id="test-user", token="test-token"), + ) + + assert exc_info.value.status_code == 400 + assert "litellm_params.mode must be a string" in exc_info.value.detail["error"] + ahealth_check.assert_not_called() + + +@pytest.mark.asyncio +async def test_test_model_connection_string_params_mode_is_the_probe_mode(): + with _test_connection_probe(MANTLE_CLAUDE_DEPLOYMENT) as ahealth_check: + await health_test_model_connection( + request=MagicMock(), + mode=None, + litellm_params={"model": "bedrock_mantle/anthropic.claude-haiku-4-5", "mode": "chat"}, + model_info={"id": "mantle-claude-id"}, + user_api_key_dict=UserAPIKeyAuth(user_id="test-user", token="test-token"), + ) + + assert ahealth_check.call_args.kwargs["mode"] == "chat" + assert "mode" not in ahealth_check.call_args.kwargs["model_params"] + + +@pytest.mark.asyncio +async def test_test_model_connection_request_mode_wins_over_resolved_mode(): + with _test_connection_probe(MANTLE_CLAUDE_DEPLOYMENT) as ahealth_check: + await health_test_model_connection( + request=MagicMock(), + mode="chat", + litellm_params={"model": "bedrock_mantle/anthropic.claude-haiku-4-5"}, + model_info={"id": "mantle-claude-id"}, + user_api_key_dict=UserAPIKeyAuth(user_id="test-user", token="test-token"), + ) + + assert ahealth_check.call_args.kwargs["mode"] == "chat" + + @pytest.mark.asyncio async def test_test_model_connection_uses_loaded_deployment_team_id(): """ diff --git a/tests/unit/proxy/test_health_check_max_tokens.py b/tests/unit/proxy/test_health_check_max_tokens.py index 33fc4cad659..e3641ac2c81 100644 --- a/tests/unit/proxy/test_health_check_max_tokens.py +++ b/tests/unit/proxy/test_health_check_max_tokens.py @@ -11,7 +11,7 @@ from litellm.proxy import health_check as hc_module from litellm.proxy.health_check import ( _is_strategy_router_deployment, _resolve_health_check_max_tokens, - _resolve_health_check_mode, + resolve_health_check_mode, _update_litellm_params_for_health_check, ) @@ -406,7 +406,7 @@ def test_update_litellm_params_health_check_reasoning_effort(): ) def test_bedrock_embedding_without_explicit_mode_skips_max_tokens(deployment_model, expected_request_model): """Embedding mode auto-detected from model cost map -> no max_tokens, provider pinned.""" - assert _resolve_health_check_mode({}, {"model": deployment_model}) == "embedding" + assert resolve_health_check_mode({}, {"model": deployment_model}) == "embedding" updated = _update_litellm_params_for_health_check({}, {"model": deployment_model}) @@ -417,12 +417,12 @@ def test_bedrock_embedding_without_explicit_mode_skips_max_tokens(deployment_mod def test_resolve_health_check_mode_prefers_explicit_model_info_mode(): """An operator-set mode wins over model-cost lookup.""" - assert _resolve_health_check_mode({"mode": "chat"}, {"model": "bedrock/amazon.titan-embed-text-v2:0"}) == "chat" + assert resolve_health_check_mode({"mode": "chat"}, {"model": "bedrock/amazon.titan-embed-text-v2:0"}) == "chat" def test_resolve_health_check_mode_unknown_model_returns_none(): - assert _resolve_health_check_mode({}, {"model": "bedrock/not-a-real-model-xyz"}) is None - assert _resolve_health_check_mode({}, {}) is None + assert resolve_health_check_mode({}, {"model": "bedrock/not-a-real-model-xyz"}) is None + assert resolve_health_check_mode({}, {}) is None def test_bedrock_chat_without_mode_still_injects_max_tokens_and_pins_provider(): @@ -481,6 +481,113 @@ async def test_run_model_health_check_threads_resolved_mode_to_ahealth_check(): assert probed_params["model"] == "amazon.titan-embed-text-v2:0" +_MANTLE_CLAUDE_DEPLOYMENT_PARAMS = { + "model": "bedrock_mantle/anthropic.claude-haiku-4-5", + "api_key": "test-bearer", + "aws_region_name": "us-east-2", +} + + +def _mantle_anthropic_response() -> dict[str, object]: + return { + "id": "msg_health", + "type": "message", + "role": "assistant", + "model": "anthropic.claude-haiku-4-5", + "content": [{"type": "text", "text": "pong"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 3, "output_tokens": 1}, + } + + +@pytest.mark.parametrize( + "deployment_model", + ["bedrock_mantle/anthropic.claude-haiku-4-5", "bedrock_mantle/anthropic.claude-opus-5-5"], +) +def test_mantle_claude_without_mode_resolves_to_anthropic_messages(deployment_model): + """Mantle only serves Claude over /anthropic/v1/messages, so that is the probe surface by default.""" + assert resolve_health_check_mode({}, {"model": deployment_model}) == "anthropic_messages" + + updated = _update_litellm_params_for_health_check({}, {"model": deployment_model}) + + assert updated["max_tokens"] == 16 + assert [message["role"] for message in updated["messages"]] == ["user"] + + +def test_mantle_claude_with_explicit_provider_param_resolves_to_anthropic_messages(): + assert ( + resolve_health_check_mode({}, {"model": "anthropic.claude-haiku-4-5", "custom_llm_provider": "bedrock_mantle"}) + == "anthropic_messages" + ) + + +def test_mantle_claude_explicit_chat_mode_wins_over_the_native_default(): + assert resolve_health_check_mode({"mode": "chat"}, {"model": "bedrock_mantle/anthropic.claude-haiku-4-5"}) == "chat" + + +@pytest.mark.parametrize( + "deployment_model", + [ + "bedrock_mantle/openai.gpt-oss-120b", + "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0", + "anthropic/claude-haiku-4-5", + ], +) +def test_native_messages_default_is_scoped_to_mantle_claude(deployment_model): + """Non-Claude Mantle ids and Claude on other providers keep their chat-completions probe.""" + assert resolve_health_check_mode({}, {"model": deployment_model}) == "chat" + + +@pytest.mark.asyncio +async def test_run_model_health_check_probes_mantle_claude_over_messages(monkeypatch): + """The deployment the ticket describes, probed end to end through the proxy's health runner. + + Before the fix the probe went to /v1/chat/completions, which Mantle answers with a + validation_error for Claude ids, so every such deployment showed unhealthy. + """ + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + + with respx.mock(assert_all_called=False) as respx_mock: + messages_route = respx_mock.post("https://bedrock-mantle.us-east-2.api.aws/anthropic/v1/messages").respond( + json=_mantle_anthropic_response() + ) + chat_route = respx_mock.post("https://bedrock-mantle.us-east-2.api.aws/v1/chat/completions").respond( + status_code=400, json={"type": "error", "error": {"type": "validation_error"}} + ) + result = await hc_module._run_model_health_check( + {"litellm_params": dict(_MANTLE_CLAUDE_DEPLOYMENT_PARAMS), "model_info": {}} + ) + + assert "error" not in result, result + assert chat_route.call_count == 0 + assert messages_route.call_count == 1 + sent = messages_route.calls.last.request + assert sent.headers["authorization"] == "Bearer test-bearer" + body = json.loads(sent.content) + assert body["model"] == "anthropic.claude-haiku-4-5" + assert body["max_tokens"] == 16 + assert [message["role"] for message in body["messages"]] == ["user"] + + +@pytest.mark.asyncio +async def test_run_model_health_check_honors_an_explicit_chat_mode_on_mantle_claude(monkeypatch): + """Negative control: an operator who pins mode=chat still gets the chat completions probe. + + Since #43646 Mantle serves Claude chat completions over its Messages endpoint as well, so + the wire no longer tells the two probes apart and the probe mode is read off the health call. + """ + fake_ahealth_check = AsyncMock(return_value={}) + monkeypatch.setattr(litellm, "ahealth_check", fake_ahealth_check) + + await hc_module._run_model_health_check( + {"litellm_params": dict(_MANTLE_CLAUDE_DEPLOYMENT_PARAMS), "model_info": {"mode": "chat"}} + ) + + assert fake_ahealth_check.call_args.kwargs["mode"] == "chat" + + def test_autodetected_embedding_skips_reasoning_effort(): """reasoning_effort must not leak into an embedding probe whose mode is auto-detected. diff --git a/ui/litellm-dashboard/src/components/add_model/add_model_modes.tsx b/ui/litellm-dashboard/src/components/add_model/add_model_modes.tsx index 81d89cc41fd..05da3f96109 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_model_modes.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_model_modes.tsx @@ -11,6 +11,7 @@ export const TEST_MODES = [ { value: "rerank", label: "Rerank - /rerank" }, { value: "realtime", label: "Realtime - /realtime" }, { value: "batch", label: "Batch - /batch" }, + { value: "anthropic_messages", label: "Anthropic Messages - /v1/messages" }, { value: "ocr", label: "OCR - /ocr" }, ]; diff --git a/ui/litellm-dashboard/src/components/model_info_view.test.tsx b/ui/litellm-dashboard/src/components/model_info_view.test.tsx index 0f29c1ef596..037ebc4040e 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.test.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.test.tsx @@ -295,6 +295,35 @@ describe("ModelInfoView", () => { expect(modelInfoArg.id).toBe("123"); }); + it("does not echo the displayed mode into the test connection request", async () => { + // /model/info fills model_info.mode in from the cost map for display. Sending that + // value back would pin the probe to it and skip the mode the provider requires, + // so the page forwards only the row's id and lets the proxy resolve the mode. + const user = userEvent.setup(); + const displayedModel = { + ...defaultModelData, + litellm_params: { ...defaultModelData.litellm_params, model: "bedrock_mantle/anthropic.claude-haiku-4-5" }, + model_info: { ...defaultModelData.model_info, mode: "chat", key: "anthropic.claude-haiku-4-5" }, + }; + mockUseModelsInfo.mockReturnValue({ data: { data: [displayedModel] }, isLoading: false, error: null }); + mockModelInfoV1Call.mockResolvedValue({ data: [displayedModel] }); + render(, { wrapper }); + + await waitFor(() => { + expect(screen.getByText("Model Settings")).toBeInTheDocument(); + }); + + await user.click(screen.getByRole("button", { name: /test connection/i })); + + await waitFor(() => { + expect(mockTestConnectionRequest).toHaveBeenCalled(); + }); + + const [, , modelInfoArg, modeArg] = mockTestConnectionRequest.mock.calls[0]; + expect(modelInfoArg).toEqual({ id: "123" }); + expect(modeArg).toBeUndefined(); + }); + it("should display error notification when connection test fails", async () => { const user = userEvent.setup(); mockTestConnectionRequest.mockRejectedValue(new Error("Connection failed")); diff --git a/ui/litellm-dashboard/src/components/model_info_view.tsx b/ui/litellm-dashboard/src/components/model_info_view.tsx index 9b33e5b10be..4fe49a28936 100644 --- a/ui/litellm-dashboard/src/components/model_info_view.tsx +++ b/ui/litellm-dashboard/src/components/model_info_view.tsx @@ -491,9 +491,7 @@ export default function ModelInfoView({ // backend silently falls back to deployments[0] and probes // the wrong endpoint. id: localModelData.model_info?.id, - mode: localModelData.model_info?.mode, }, - localModelData.model_info?.mode, ); if (response.status === "success") { diff --git a/ui/litellm-dashboard/src/components/networking.tsx b/ui/litellm-dashboard/src/components/networking.tsx index 52021f61ba6..4e606222eae 100644 --- a/ui/litellm-dashboard/src/components/networking.tsx +++ b/ui/litellm-dashboard/src/components/networking.tsx @@ -2165,7 +2165,7 @@ export const testConnectionRequest = async ( accessToken: string, litellm_params: Record, model_info: Record, - mode: string, + mode?: string, ) => { try { // Construct the URL based on environment diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index f0423a8215b..d1ef1e92008 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -27407,9 +27407,9 @@ export interface components { }; /** * Mode - * @description The mode to test the model with. If not provided, auto-detected from model capabilities. + * @description The mode to test the model with. If not provided, resolved the way /health does: the deployment's model_info.mode (only while the request tests the deployment's own model), then the mode the provider requires for that model, then the model cost map. */ - mode?: ("chat" | "completion" | "embedding" | "audio_speech" | "audio_transcription" | "image_generation" | "image_edit" | "video_generation" | "batch" | "rerank" | "realtime" | "responses" | "ocr") | null; + mode?: ("chat" | "completion" | "embedding" | "audio_speech" | "audio_transcription" | "image_generation" | "image_edit" | "video_generation" | "batch" | "rerank" | "realtime" | "responses" | "anthropic_messages" | "ocr") | null; /** * Model Info * @description Model info for the health check From dfdd496db8e1954993e4a2096cbb450c66db10db Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 17:07:10 -0700 Subject: [PATCH 10/18] fix(tests): match the OS bind error in the owned-proxy port-race retry (#44462) * fix(tests): match the OS bind error in the owned-proxy port-race retry * test(integration): keep the port-race predicate pure so its unit tests stay in-process --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> --- tests/integration/_support/process.py | 11 ++++-- .../unit/integration_support/test_process.py | 38 +++++++++++++++++++ 2 files changed, 45 insertions(+), 4 deletions(-) create mode 100644 tests/unit/integration_support/test_process.py diff --git a/tests/integration/_support/process.py b/tests/integration/_support/process.py index 0d199770181..3d76c7881e1 100644 --- a/tests/integration/_support/process.py +++ b/tests/integration/_support/process.py @@ -1,3 +1,4 @@ +import errno import os import signal import socket @@ -112,6 +113,7 @@ def _stop(process: subprocess.Popen[bytes]) -> None: _PORT_ATTEMPTS: Final = 3 +_BIND_COLLISION: Final = os.strerror(errno.EADDRINUSE) def _free_port() -> int: @@ -142,8 +144,8 @@ def _launch(command: tuple[str, ...], root: Path, environment: Mapping[str, str] return _Launch(process, port, log_path) -def _lost_port_race(launch: _Launch) -> bool: - return launch.process.poll() is not None and "address already in use" in launch.log.read_text() +def _lost_port_race(exit_code: int | None, log: Path) -> bool: + return exit_code is not None and _BIND_COLLISION in log.read_text() def _wait_until_ready(launch: _Launch) -> None: @@ -165,13 +167,14 @@ def _launch_until_bound( launch: Final = _launch(command, root, environment, output) try: _wait_until_ready(launch) - assert launch.process.poll() is None or (attempts > 1 and _lost_port_race(launch)), ( + exit_code: Final = launch.process.poll() + assert exit_code is None or (attempts > 1 and _lost_port_race(exit_code, launch.log)), ( "Owned proxy exited before readiness" ) except BaseException: _stop(launch.process) raise - if launch.process.poll() is None: + if exit_code is None: return launch _stop(launch.process) return _launch_until_bound(command, root, environment, output, attempts - 1) diff --git a/tests/unit/integration_support/test_process.py b/tests/unit/integration_support/test_process.py new file mode 100644 index 00000000000..042a2447dff --- /dev/null +++ b/tests/unit/integration_support/test_process.py @@ -0,0 +1,38 @@ +from __future__ import annotations + +import errno +import importlib +import os +from pathlib import Path +from types import ModuleType +from typing import Final + +import pytest + +TESTS_DIR: Final = Path(__file__).resolve().parents[2] +BIND_ERROR_LINE: Final = f"ERROR: {OSError(errno.EADDRINUSE, os.strerror(errno.EADDRINUSE))}\n" +UNRELATED_CRASH: Final = "Traceback (most recent call last):\nModuleNotFoundError: No module named 'litellm'\n" + + +@pytest.fixture +def process_module(monkeypatch: pytest.MonkeyPatch) -> ModuleType: + monkeypatch.syspath_prepend(str(TESTS_DIR)) + return importlib.import_module("integration._support.process") + + +def _written_log(directory: Path, text: str) -> Path: + log: Final = directory / "owned-proxy.log" + log.write_text(text) + return log + + +def test_lost_port_race_matches_the_bind_error_the_server_logs(process_module: ModuleType, tmp_path: Path) -> None: + assert process_module._lost_port_race(1, _written_log(tmp_path, BIND_ERROR_LINE)) + + +def test_lost_port_race_ignores_an_exit_for_another_reason(process_module: ModuleType, tmp_path: Path) -> None: + assert not process_module._lost_port_race(1, _written_log(tmp_path, UNRELATED_CRASH)) + + +def test_lost_port_race_needs_the_process_to_have_exited(process_module: ModuleType, tmp_path: Path) -> None: + assert not process_module._lost_port_race(None, _written_log(tmp_path, BIND_ERROR_LINE)) From 0ed1c08f024791ebeecdc09ddd4c792a62f2a65a Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sat, 3 Oct 2026 17:08:30 -0700 Subject: [PATCH 11/18] feat(anthropic): workload identity federation and pluggable identity sources (#44448) * feat(anthropic): workload identity federation and pluggable identity sources Backend half of #38818 (internal copy of the fork PR #38013), rebuilt as one commit on top of litellm_internal_staging without the dashboard changes. Deployments on anthropic/ without a static api_key can exchange an OIDC workload assertion for a short-lived sk-ant-oat01 token through a shared RFC 7523 JWT-bearer engine. The assertion comes from a mounted token file, an env token, a LiteLLM-signed issuer, or Keycloak, chosen per deployment, per named credential, or through ANTHROPIC_IDENTITY_SOURCE. The federation fields are server-owned: refused inline in request bodies and on POST /model/new, proxy-admin only on credentials, and the token exchange is pinned to api.anthropic.com unless LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS adds a host. GET /credentials/{name}/jwks exports the public key set of a LiteLLM-signed credential for the Claude Console. The OpenAI federation trio from #39613 rides along on the backend side with the same server-owned handling. Fixes #28607 Resolves LIT-6107 Co-authored-by: derhornspieler <15236687+derhornspieler@users.noreply.github.com> * fix(anthropic): let batch-result downloads mint from deployment params and accept host:port allowlist entries The files handler enabled workload identity on batch-result downloads but never received the deployment's litellm_params, so a deployment authenticating through a named credential could only mint from process-wide env vars. It now threads litellm_params through to the auth header the way the batch retrieve path already does. LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS entries written as host:port were read by urlsplit as a scheme, so the allowlist kept the raw entry while the exchange compared bare hostnames and refused the gateway. Entries are now parsed as network locations whether or not they carry a scheme. * fix(types): move the WIF kwargs key sets to a leaf module so the kwargs funnel imports without a cycle * test(anthropic): pin case-insensitive matching of WIF exchange-host allowlist entries * fix(anthropic): end workload identity federation errors without a period so the router suffix reads cleanly * fix(proxy): decrypt stored litellm_params before the WIF write gate * fix(proxy): hide WIF secret references from /health output * fix(proxy): keep the proxy error shape on credential endpoint refusals * fix(proxy): hide identity token file paths from /health output * fix(anthropic): rename the federation workspace param so Bedrock's anthropic_workspace_id keeps working The Bedrock Claude Platform route already reads anthropic_workspace_id from optional_params, so banning that spelling as a server-owned federation parameter broke a pre-existing client capability. The federation field is now anthropic_federation_workspace_id (env ANTHROPIC_FEDERATION_WORKSPACE_ID), which restores the base branch's behavior for Bedrock callers, drops the Bedrock-specific hint from the refusal message, and deletes the unconditional ban constant that no longer had a reader * fix(auth): share one exchanged token across workers reading the same assertion Anthropic accepts each identity assertion exactly once, so two uvicorn workers reading the same token file both minting from it means the second exchange is denied with jti_reused. Minted tokens now land in a per-user 0700 cache directory guarded by a file lock, so workers on the same host reuse one exchange until the token expires or the assertion rotates. A 401 is only retried when the re-read assertion actually differs, and the denial hint explains jti_reused. LITELLM_TOKEN_EXCHANGE_CACHE_DIR moves the cache and an empty value disables it * fix: keep anthropic federation from being shadowed or leaked An empty or whitespace-only ANTHROPIC_API_KEY counted as set, so a federated deployment sent an empty x-api-key on every call instead of minting a token. Blank values now read as unset, and a real static key on a federated deployment logs once that it outranks federation and nothing is being federated. The exchange-host allowlist matched hostnames only, so a second process on another port of an allowed host was trusted with the workload's identity token. An entry that names a port now trusts that port alone, while a bare host still trusts every port. The shared token store exists so the workers reading one projected token file do not each spend its single-use jti. A source that mints its own assertion per exchange shares nothing with another worker, so it no longer writes a live token to disk for a lookup that can never hit. * fix: unlink a staged token file a failed write leaves behind The 401 denial hint now also says federation ignores ANTHROPIC_WORKSPACE_ID, which the Bedrock Claude platform provider already reads. * refactor: move anthropic jwks derivation behind a provider-owned tagged union * fix: unlink the staged token file when its write fails at close A buffered write only reaches the disk when the handle closes, so a full disk surfaces at close and left the staging file behind holding a usable token. * fix(anthropic): close the staging descriptor before writing the shared token file * fix(wif): judge federation writes by what they set, not what is stored The admin gate read the stored deployment, so a team admin lost edit, delete and Test Connection on any deployment carrying federation params. It now returns early unless the submitted fields touch the federation surface, and a Test Connection probe that points the deployment at its own api_base is still refused, with the 403 no longer wrapped into a 500 The rest of the same review pass: POST /model/new refuses only a blocking value of `blocked`, so a client that always sends `blocked: false` is not turned away; a request body can no longer pick which federated identity to mint as by naming a stored credential; an advisory refresh the executor refuses disarms the entry instead of wedging the identity until the follower timeout; the static-key shadow warning resolves its env fallback inside the cache instead of once per request; credential writes drop nulls before storing them; the token exchange validates the endpoint URL before reading an assertion and keeps refusing redirects across a client heal; /health hides every server-owned federation field from non-admins; and the async create_file and create_batch paths say which setting is missing when the provider resolves no URL * fix(proxy): let a deployment write name a federated credential reject_federated_credential_reference runs from is_request_body_safe, which pre_db_read_auth_checks calls on every route, so it also fired on POST /model/new, /model/update, /model/{id}/update and /health/test_connection. A proxy admin could no longer attach a federated credential to a deployment over the API or the Admin UI, leaving a static config.yaml entry as the only way to configure the feature the rejection told the caller to go configure, and _reject_non_admin_wif_write never got to make the call it exists to make. is_request_body_safe now takes the route and skips only the credential-reference check on the routes that reach can_user_make_model_call. Federation fields typed inline into a body stay refused everywhere, and a call naming a federated credential still cannot pick the identity it mints as. * refactor(proxy): derive health display policy from the federation key sets The health check module hand-copied the five workload identity fields whose value is a credential, so a shared proxy surface named provider-specific parameters and a newly added secret-bearing field would have gone on being displayed until someone remembered both places WIF_SECRET_BEARING_KEYS now sits beside the key sets it splits out of, types/utils derives secret_bearing_wif_litellm_params from it, and the health layer splats that tuple the same way it already splats the admin-only one * fix(anthropic_wif): treat blank identity-source fields as unset * test(proxy): classify the federation params in the credential slot registry main's registry test (#43298) now fails the build for any credential-named deployment param without a classification. The five federation fields that carry a token, a token file path, or a signing or client secret reference are Unplanted, matching WIF_SECRET_BEARING_KEYS; the four remaining Keycloak settings name a URL, a client id, an auth method, or a scope and are NotSecret * fix(anthropic_wif): declare federation params as owned connection leaves and chart their metrics Register the 18 Anthropic and 3 OpenAI federation params as frozen ConnectionSettings leaves so the owned-kwarg registry, the kwargs funnel and the request-body ban list read one declaration. Pass the deployment api_base through to the count-tokens handler instead of a pre-suffixed URL, which doubled the /count_tokens path on main's prompt-cache predictor. Add the five litellm_anthropic_wif_* families to the all-metrics Grafana dashboard. * fix(credentials): gate PATCH on WIF fields resolved from model_id The credential PATCH handler checked server-owned workload identity federation fields only on the values the caller sent, while a body that named a deployment through model_id had its credential values resolved after that check. A non-admin could therefore copy a federated deployment's WIF fields onto an ordinary credential. Resolve the incoming values first and run the non-admin gate on them, matching the POST path * fix(anthropic): count tokens with ANTHROPIC_AUTH_TOKEN through the shared auth header Count-tokens walked its own credential ladder: a static key, else skip minting when ANTHROPIC_AUTH_TOKEN is set, else mint a federated token. With only the auth token set it forwarded nothing and the proxy silently fell back to its local tokenizer while chat on the same deployment authenticated with that token. The handler now takes the auth header that AnthropicModelInfo.aget_auth_header resolves, the same ladder chat, files, batches and skills use, and merges the oauth beta a minted or consumer token carries with the token-counting beta --------- Co-authored-by: mateo-berri <277851410+mateo-berri@users.noreply.github.com> Co-authored-by: derhornspieler <15236687+derhornspieler@users.noreply.github.com> Co-authored-by: mateo-berri --- .../grafana_dashboard.json | 127 +- .../dashboard_all_metrics/readme.md | 2 +- litellm/batches/batch_utils.py | 4 + litellm/batches/main.py | 12 +- .../litellm_core_utils/get_litellm_params.py | 3 + litellm/llms/anthropic/batches/handler.py | 18 +- .../llms/anthropic/batches/transformation.py | 28 +- litellm/llms/anthropic/chat/transformation.py | 4 +- litellm/llms/anthropic/common_utils.py | 353 +++- .../llms/anthropic/count_tokens/handler.py | 10 +- .../anthropic/count_tokens/token_counter.py | 31 +- .../anthropic/count_tokens/transformation.py | 42 +- litellm/llms/anthropic/files/handler.py | 13 +- .../llms/anthropic/files/transformation.py | 62 +- .../pass_through/messages/transformation.py | 124 +- .../llms/anthropic/prompt_cache_prediction.py | 8 +- .../llms/anthropic/skills/transformation.py | 56 +- litellm/llms/anthropic/wif.py | 604 ++++++ litellm/llms/azure_ai/embed/handler.py | 2 + .../anthropic_messages/transformation.py | 23 + litellm/llms/base_llm/auth/__init__.py | 99 + .../llms/base_llm/auth/client_credentials.py | 225 ++ litellm/llms/base_llm/auth/identity_source.py | 76 + litellm/llms/base_llm/auth/internal_issuer.py | 86 + litellm/llms/base_llm/auth/jwt_signing.py | 115 ++ .../llms/base_llm/auth/shared_token_store.py | 181 ++ litellm/llms/base_llm/auth/token_exchange.py | 941 +++++++++ litellm/llms/base_llm/auth/types.py | 100 + litellm/llms/base_llm/base_utils.py | 17 + litellm/llms/custom_httpx/http_handler.py | 6 +- litellm/llms/custom_httpx/llm_http_handler.py | 264 ++- .../llms/openai/chat/gpt_transformation.py | 53 +- litellm/llms/openai/openai.py | 16 +- .../llms/openai/responses/transformation.py | 7 +- litellm/llms/openai/workload_identity.py | 19 +- litellm/main.py | 17 +- litellm/models/credentials.py | 7 +- litellm/proxy/auth/auth_utils.py | 90 +- .../common_utils/credential_hydration.py | 189 ++ .../proxy/credential_endpoints/endpoints.py | 232 ++- litellm/proxy/health_check.py | 14 +- .../health_endpoints/_health_endpoints.py | 16 +- .../model_management_endpoints.py | 102 +- .../llm_passthrough_endpoints.py | 23 +- litellm/proxy/proxy_server.py | 20 +- litellm/router.py | 22 +- .../clientside_credential_handler.py | 17 + .../router_utils/fallback_event_handlers.py | 10 +- litellm/secret_managers/main.py | 8 +- litellm/types/litellm_params.py | 31 + litellm/types/llms/anthropic.py | 1 + litellm/types/router.py | 91 +- litellm/types/services.py | 5 + litellm/types/utils.py | 10 + litellm/types/workload_identity.py | 35 + litellm/utils.py | 5 +- .../endpointaudit/coverage_allowlist.txt | 1 + tests/unit/batches/test_batch_utils.py | 34 + .../integrations/test_prometheus_services.py | 25 + .../test_get_litellm_params.py | 145 ++ .../llms/anthropic/batches/test_handler.py | 146 +- .../anthropic/batches/test_transformation.py | 47 +- .../test_anthropic_guardrail_handler.py | 46 +- .../chat/test_anthropic_chat_handler.py | 321 +-- ...est_code_interpreter_results_extraction.py | 18 +- .../test_anthropic_files_transformation.py | 141 +- .../messages/test_advisor_orchestration.py | 24 +- .../test_handler_output_config_passthrough.py | 8 +- .../test_streaming_iterator_combined_chunk.py | 26 +- .../test_streaming_iterator_compaction.py | 34 +- .../test_streaming_iterator_empty_choices.py | 8 +- .../test_streaming_iterator_first_delta.py | 14 +- .../test_streaming_iterator_tool_args.py | 42 +- .../test_clear_tool_uses.py | 8 +- .../context_management/test_compact.py | 73 +- .../context_management/test_dispatcher.py | 4 +- .../messages/test_advisor_integration.py | 27 +- .../test_agentic_streaming_iterator.py | 52 +- ...al_pass_through_messages_transformation.py | 118 ++ .../messages/test_anthropic_messages_speed.py | 12 +- ...t_anthropic_messages_structured_outputs.py | 4 +- .../test_content_after_stop_reason.py | 82 +- .../messages/test_parallel_tool_calls.py | 26 +- .../test_reasoning_auto_summary_messages.py | 25 +- .../test_request_optional_param_utils.py | 33 +- .../pass_through/messages/test_sse_wrapper.py | 43 +- .../messages/test_streaming_iterator.py | 7 +- .../anthropic/test_anthropic_common_utils.py | 1625 ++++++++++++++- ...t_anthropic_count_tokens_transformation.py | 74 + .../test_anthropic_files_and_batches.py | 296 ++- .../test_anthropic_prompt_cache_prediction.py | 15 +- .../unit/llms/anthropic/test_anthropic_wif.py | 1293 ++++++++++++ .../test_cost_calculation_dict_safety.py | 10 +- .../llms/anthropic/test_count_tokens_oauth.py | 249 ++- .../anthropic/test_message_sanitization.py | 65 +- tests/unit/llms/base_llm/auth/__init__.py | 0 .../base_llm/auth/test_client_credentials.py | 484 +++++ .../base_llm/auth/test_identity_source.py | 239 +++ .../base_llm/auth/test_internal_issuer.py | 188 ++ .../llms/base_llm/auth/test_jwt_signing.py | 215 ++ .../base_llm/auth/test_shared_token_store.py | 321 +++ .../llms/base_llm/auth/test_token_exchange.py | 1807 +++++++++++++++++ .../llms/custom_httpx/test_http_handler.py | 47 + .../custom_httpx/test_llm_http_handler.py | 281 ++- .../openai/test_openai_workload_identity.py | 351 ++++ tests/unit/models/test_models.py | 38 + tests/unit/proxy/auth/test_auth_utils.py | 400 ++-- .../common_utils/test_credential_hydration.py | 28 + .../credential_endpoints/test_endpoints.py | 1091 +++++++++- .../health_endpoints/test_health_endpoints.py | 246 ++- .../test_model_management_endpoints.py | 1348 +++++++++--- .../test_llm_pass_through_endpoints.py | 806 +++++--- .../proxy/proxy_server/test_proxy_config.py | 17 + .../proxy/test_credential_slot_registry.py | 9 + ...test_proxy_server_endpoints_and_startup.py | 51 + .../test_fallback_event_handlers.py | 56 +- .../test_anthropic_skills_transformation.py | 77 +- tests/unit/test_lazy_imports.py | 12 + tests/unit/test_router/test_router.py | 59 + tests/unit/types/test_litellm_params.py | 50 + tests/unit/types/test_router.py | 113 +- ui/litellm-dashboard/src/lib/http/schema.d.ts | 142 ++ 122 files changed, 16215 insertions(+), 2158 deletions(-) create mode 100644 litellm/llms/anthropic/wif.py create mode 100644 litellm/llms/base_llm/auth/__init__.py create mode 100644 litellm/llms/base_llm/auth/client_credentials.py create mode 100644 litellm/llms/base_llm/auth/identity_source.py create mode 100644 litellm/llms/base_llm/auth/internal_issuer.py create mode 100644 litellm/llms/base_llm/auth/jwt_signing.py create mode 100644 litellm/llms/base_llm/auth/shared_token_store.py create mode 100644 litellm/llms/base_llm/auth/token_exchange.py create mode 100644 litellm/llms/base_llm/auth/types.py create mode 100644 litellm/proxy/common_utils/credential_hydration.py create mode 100644 litellm/types/workload_identity.py create mode 100644 tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_transformation.py create mode 100644 tests/unit/llms/anthropic/test_anthropic_wif.py create mode 100644 tests/unit/llms/base_llm/auth/__init__.py create mode 100644 tests/unit/llms/base_llm/auth/test_client_credentials.py create mode 100644 tests/unit/llms/base_llm/auth/test_identity_source.py create mode 100644 tests/unit/llms/base_llm/auth/test_internal_issuer.py create mode 100644 tests/unit/llms/base_llm/auth/test_jwt_signing.py create mode 100644 tests/unit/llms/base_llm/auth/test_shared_token_store.py create mode 100644 tests/unit/llms/base_llm/auth/test_token_exchange.py create mode 100644 tests/unit/proxy/common_utils/test_credential_hydration.py diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json index 5bd7ed97a55..4fb926a9658 100644 --- a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/grafana_dashboard.json @@ -5710,6 +5710,17 @@ } }, "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "histogram_quantile(0.95, sum(rate(litellm_anthropic_wif_latency_bucket[$__rate_interval])) by (le))", + "legendFormat": "anthropic_wif", + "range": true, + "refId": "A" + }, { "datasource": { "type": "prometheus", @@ -5719,7 +5730,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_auth_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "auth", "range": true, - "refId": "A" + "refId": "B" }, { "datasource": { @@ -5730,7 +5741,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_batch_write_to_db_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "batch_write_to_db", "range": true, - "refId": "B" + "refId": "C" }, { "datasource": { @@ -5741,7 +5752,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_postgres_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "postgres", "range": true, - "refId": "C" + "refId": "D" }, { "datasource": { @@ -5752,7 +5763,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_proxy_pre_call_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "proxy_pre_call", "range": true, - "refId": "D" + "refId": "E" }, { "datasource": { @@ -5763,7 +5774,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_redis_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "redis", "range": true, - "refId": "E" + "refId": "F" }, { "datasource": { @@ -5774,7 +5785,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_redis_daily_org_spend_update_queue_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "redis_daily_org_spend_update_queue", "range": true, - "refId": "F" + "refId": "G" }, { "datasource": { @@ -5785,7 +5796,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_redis_daily_tag_spend_update_queue_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "redis_daily_tag_spend_update_queue", "range": true, - "refId": "G" + "refId": "H" }, { "datasource": { @@ -5796,7 +5807,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_redis_daily_team_spend_update_queue_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "redis_daily_team_spend_update_queue", "range": true, - "refId": "H" + "refId": "I" }, { "datasource": { @@ -5807,7 +5818,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_redis_window_spend_update_queue_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "redis_window_spend_update_queue", "range": true, - "refId": "I" + "refId": "J" }, { "datasource": { @@ -5818,7 +5829,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_reset_budget_job_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "reset_budget_job", "range": true, - "refId": "J" + "refId": "K" }, { "datasource": { @@ -5829,7 +5840,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_router_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "router", "range": true, - "refId": "K" + "refId": "L" }, { "datasource": { @@ -5840,7 +5851,7 @@ "expr": "histogram_quantile(0.95, sum(rate(litellm_self_latency_bucket[$__rate_interval])) by (le))", "legendFormat": "self", "range": true, - "refId": "L" + "refId": "M" } ], "title": "Service latency p95 (litellm__latency)", @@ -5888,6 +5899,28 @@ } }, "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "sum(rate(litellm_anthropic_wif_total_requests_total[$__rate_interval]))", + "legendFormat": "anthropic_wif", + "range": true, + "refId": "A" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "sum(rate(litellm_anthropic_wif_cache_total_requests_total[$__rate_interval]))", + "legendFormat": "anthropic_wif_cache", + "range": true, + "refId": "B" + }, { "datasource": { "type": "prometheus", @@ -5897,7 +5930,7 @@ "expr": "sum(rate(litellm_auth_total_requests_total[$__rate_interval]))", "legendFormat": "auth", "range": true, - "refId": "A" + "refId": "C" }, { "datasource": { @@ -5908,7 +5941,7 @@ "expr": "sum(rate(litellm_batch_write_to_db_total_requests_total[$__rate_interval]))", "legendFormat": "batch_write_to_db", "range": true, - "refId": "B" + "refId": "D" }, { "datasource": { @@ -5919,7 +5952,7 @@ "expr": "sum(rate(litellm_postgres_total_requests_total[$__rate_interval]))", "legendFormat": "postgres", "range": true, - "refId": "C" + "refId": "E" }, { "datasource": { @@ -5930,7 +5963,7 @@ "expr": "sum(rate(litellm_proxy_pre_call_total_requests_total[$__rate_interval]))", "legendFormat": "proxy_pre_call", "range": true, - "refId": "D" + "refId": "F" }, { "datasource": { @@ -5941,7 +5974,7 @@ "expr": "sum(rate(litellm_redis_total_requests_total[$__rate_interval]))", "legendFormat": "redis", "range": true, - "refId": "E" + "refId": "G" }, { "datasource": { @@ -5952,7 +5985,7 @@ "expr": "sum(rate(litellm_redis_daily_org_spend_update_queue_total_requests_total[$__rate_interval]))", "legendFormat": "redis_daily_org_spend_update_queue", "range": true, - "refId": "F" + "refId": "H" }, { "datasource": { @@ -5963,7 +5996,7 @@ "expr": "sum(rate(litellm_redis_daily_tag_spend_update_queue_total_requests_total[$__rate_interval]))", "legendFormat": "redis_daily_tag_spend_update_queue", "range": true, - "refId": "G" + "refId": "I" }, { "datasource": { @@ -5974,7 +6007,7 @@ "expr": "sum(rate(litellm_redis_daily_team_spend_update_queue_total_requests_total[$__rate_interval]))", "legendFormat": "redis_daily_team_spend_update_queue", "range": true, - "refId": "H" + "refId": "J" }, { "datasource": { @@ -5985,7 +6018,7 @@ "expr": "sum(rate(litellm_redis_window_spend_update_queue_total_requests_total[$__rate_interval]))", "legendFormat": "redis_window_spend_update_queue", "range": true, - "refId": "I" + "refId": "K" }, { "datasource": { @@ -5996,7 +6029,7 @@ "expr": "sum(rate(litellm_reset_budget_job_total_requests_total[$__rate_interval]))", "legendFormat": "reset_budget_job", "range": true, - "refId": "J" + "refId": "L" }, { "datasource": { @@ -6007,7 +6040,7 @@ "expr": "sum(rate(litellm_router_total_requests_total[$__rate_interval]))", "legendFormat": "router", "range": true, - "refId": "K" + "refId": "M" }, { "datasource": { @@ -6018,7 +6051,7 @@ "expr": "sum(rate(litellm_self_total_requests_total[$__rate_interval]))", "legendFormat": "self", "range": true, - "refId": "L" + "refId": "N" } ], "title": "Service request rate (litellm__total_requests)", @@ -6066,6 +6099,28 @@ } }, "targets": [ + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "sum(rate(litellm_anthropic_wif_failed_requests_total[$__rate_interval])) by (error_class)", + "legendFormat": "anthropic_wif / {{error_class}}", + "range": true, + "refId": "A" + }, + { + "datasource": { + "type": "prometheus", + "uid": "${DS_PROMETHEUS}" + }, + "editorMode": "code", + "expr": "sum(rate(litellm_anthropic_wif_cache_failed_requests_total[$__rate_interval])) by (error_class)", + "legendFormat": "anthropic_wif_cache / {{error_class}}", + "range": true, + "refId": "B" + }, { "datasource": { "type": "prometheus", @@ -6075,7 +6130,7 @@ "expr": "sum(rate(litellm_auth_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "auth / {{error_class}}", "range": true, - "refId": "A" + "refId": "C" }, { "datasource": { @@ -6086,7 +6141,7 @@ "expr": "sum(rate(litellm_batch_write_to_db_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "batch_write_to_db / {{error_class}}", "range": true, - "refId": "B" + "refId": "D" }, { "datasource": { @@ -6097,7 +6152,7 @@ "expr": "sum(rate(litellm_postgres_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "postgres / {{error_class}}", "range": true, - "refId": "C" + "refId": "E" }, { "datasource": { @@ -6108,7 +6163,7 @@ "expr": "sum(rate(litellm_proxy_pre_call_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "proxy_pre_call / {{error_class}}", "range": true, - "refId": "D" + "refId": "F" }, { "datasource": { @@ -6119,7 +6174,7 @@ "expr": "sum(rate(litellm_redis_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "redis / {{error_class}}", "range": true, - "refId": "E" + "refId": "G" }, { "datasource": { @@ -6130,7 +6185,7 @@ "expr": "sum(rate(litellm_redis_daily_org_spend_update_queue_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "redis_daily_org_spend_update_queue / {{error_class}}", "range": true, - "refId": "F" + "refId": "H" }, { "datasource": { @@ -6141,7 +6196,7 @@ "expr": "sum(rate(litellm_redis_daily_tag_spend_update_queue_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "redis_daily_tag_spend_update_queue / {{error_class}}", "range": true, - "refId": "G" + "refId": "I" }, { "datasource": { @@ -6152,7 +6207,7 @@ "expr": "sum(rate(litellm_redis_daily_team_spend_update_queue_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "redis_daily_team_spend_update_queue / {{error_class}}", "range": true, - "refId": "H" + "refId": "J" }, { "datasource": { @@ -6163,7 +6218,7 @@ "expr": "sum(rate(litellm_redis_window_spend_update_queue_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "redis_window_spend_update_queue / {{error_class}}", "range": true, - "refId": "I" + "refId": "K" }, { "datasource": { @@ -6174,7 +6229,7 @@ "expr": "sum(rate(litellm_reset_budget_job_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "reset_budget_job / {{error_class}}", "range": true, - "refId": "J" + "refId": "L" }, { "datasource": { @@ -6185,7 +6240,7 @@ "expr": "sum(rate(litellm_router_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "router / {{error_class}}", "range": true, - "refId": "K" + "refId": "M" }, { "datasource": { @@ -6196,7 +6251,7 @@ "expr": "sum(rate(litellm_self_failed_requests_total[$__rate_interval])) by (error_class)", "legendFormat": "self / {{error_class}}", "range": true, - "refId": "L" + "refId": "N" } ], "title": "Service failure rate (litellm__failed_requests)", diff --git a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md index a3869213be2..5d3d2c159c1 100644 --- a/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md +++ b/cookbook/litellm_proxy_server/grafana_dashboard/dashboard_all_metrics/readme.md @@ -1,6 +1,6 @@ # LiteLLM All Prometheus Metrics dashboard -Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about +Every `litellm_*` metric family the proxy can expose on `/metrics` (141 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard diff --git a/litellm/batches/batch_utils.py b/litellm/batches/batch_utils.py index c58b0d721ad..23ef2a99585 100644 --- a/litellm/batches/batch_utils.py +++ b/litellm/batches/batch_utils.py @@ -17,6 +17,7 @@ from litellm.llms.vertex_ai.batches.transformation import ( ) from litellm.types.llms.openai import Batch from litellm.types.utils import ModelInfo, Usage +from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS from litellm.utils import token_counter @@ -543,6 +544,9 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict: "max_retries", "_litellm_internal_model_credentials", *AWS_CREDENTIAL_KWARGS_KEYS, + # A federated deployment holds no api_key, so without these the fetch that reads a + # finished batch's output has nothing to authenticate with and its cost is never billed. + *sorted(ANTHROPIC_WIF_KWARGS_KEYS), ) for key in credential_keys: if key in litellm_params: diff --git a/litellm/batches/main.py b/litellm/batches/main.py index f977fc03891..3ba27cc9791 100644 --- a/litellm/batches/main.py +++ b/litellm/batches/main.py @@ -200,7 +200,7 @@ def create_batch( LiteLLM Equivalent of POST: https://api.openai.com/v1/batches """ try: - optional_params: Final = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams.model_validate(kwargs) litellm_call_id: Final = kwargs.get("litellm_call_id", None) proxy_server_request: Final = kwargs.get("proxy_server_request", None) model_info: Final = kwargs.get("model_info", None) @@ -217,7 +217,7 @@ def create_batch( ) _is_async: Final = kwargs.pop("acreate_batch", False) is True - litellm_params: Final = dict(GenericLiteLLMParams(**kwargs)) + litellm_params: Final = dict(GenericLiteLLMParams.model_validate(kwargs)) litellm_logging_obj: Final[LiteLLMLoggingObj] = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None)) ### TIMEOUT LOGIC ### timeout: Final = _resolve_timeout(optional_params, kwargs, custom_llm_provider) @@ -530,6 +530,7 @@ def _handle_retrieve_batch_providers_without_provider_config( ) api_key = optional_params.api_key or litellm.api_key or litellm.azure_key or get_secret_str("ANTHROPIC_API_KEY") + batch_params: Final = dict(litellm_params) response = anthropic_batches_instance.retrieve_batch( _is_async=_is_async, batch_id=batch_id, @@ -537,6 +538,7 @@ def _handle_retrieve_batch_providers_without_provider_config( api_key=api_key, timeout=timeout, max_retries=optional_params.max_retries, + litellm_params=batch_params, ) else: raise litellm.exceptions.BadRequestError( @@ -573,7 +575,7 @@ def retrieve_batch( LiteLLM Equivalent of GET https://api.openai.com/v1/batches/{batch_id} """ try: - optional_params: Final = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams.model_validate(kwargs) litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None) ### TIMEOUT LOGIC ### timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600 @@ -755,7 +757,7 @@ def list_batches( """ try: # set API KEY - optional_params: Final = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams.model_validate(kwargs) litellm_params: Final = get_litellm_params( custom_llm_provider=custom_llm_provider, **kwargs, @@ -956,7 +958,7 @@ def cancel_batch( verbose_logger.exception( "litellm.batches.main.py::cancel_batch() - Error inferring custom_llm_provider - %s", e ) - optional_params: Final = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams.model_validate(kwargs) litellm_params: Final = get_litellm_params( custom_llm_provider=custom_llm_provider, **kwargs, diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 5adb9a80f9c..0fe1952db4b 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -11,6 +11,7 @@ from litellm.litellm_core_utils.core_helpers import normalize_drop_params from litellm.llms.openai.data_residency import infer_openai_data_residency from litellm.types.litellm_params import MAX_CONTROL_INT_DIGITS, ControlOptions from litellm.types.router import CustomPricingLiteLLMParams +from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS, OPENAI_WIF_KWARGS_KEYS AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset( { @@ -70,6 +71,8 @@ OPTIONAL_KWARGS_KEYS: Final = ( } ) | AWS_CREDENTIAL_KWARGS_KEYS + | ANTHROPIC_WIF_KWARGS_KEYS + | OPENAI_WIF_KWARGS_KEYS | frozenset(CustomPricingLiteLLMParams.model_fields) ) diff --git a/litellm/llms/anthropic/batches/handler.py b/litellm/llms/anthropic/batches/handler.py index f418e7e08be..3f3716e0de6 100644 --- a/litellm/llms/anthropic/batches/handler.py +++ b/litellm/llms/anthropic/batches/handler.py @@ -42,6 +42,7 @@ class AnthropicBatchesHandler: timeout: float | httpx.Timeout, max_retries: int | None, logging_obj: LiteLLMLoggingObj | None = None, + litellm_params: dict | None = None, # mutable-ok: handed straight to validate_environment ) -> LiteLLMBatch: """ Async: Retrieve a batch from Anthropic. @@ -60,9 +61,7 @@ class AnthropicBatchesHandler: # Resolve API credentials api_base = api_base or self.anthropic_model_info.get_api_base(api_base) api_key = api_key or self.anthropic_model_info.get_api_key() - - if not api_key: - raise ValueError("Missing Anthropic API Key") + resolved_litellm_params: Final = litellm_params if litellm_params is not None else {} # Create a minimal logging object if not provided if logging_obj is None: @@ -85,16 +84,18 @@ class AnthropicBatchesHandler: api_base=api_base, batch_id=batch_id, optional_params={}, - litellm_params={}, + litellm_params=resolved_litellm_params, ) - # Validate environment and get headers - headers: Final = self.provider_config.validate_environment( + # Validate environment and get headers. Offloaded to a worker thread: a WIF token + # exchange here would otherwise block the event loop. + headers: Final = await asyncio.to_thread( + self.provider_config.validate_environment, headers={}, model="", messages=[], optional_params={}, - litellm_params={}, + litellm_params=resolved_litellm_params, api_key=api_key, api_base=api_base, ) @@ -130,6 +131,7 @@ class AnthropicBatchesHandler: timeout: float | httpx.Timeout, max_retries: int | None, logging_obj: LiteLLMLoggingObj | None = None, + litellm_params: dict | None = None, # mutable-ok: handed straight to validate_environment ) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]: """ Retrieve a batch from Anthropic. @@ -154,6 +156,7 @@ class AnthropicBatchesHandler: timeout=timeout, max_retries=max_retries, logging_obj=logging_obj, + litellm_params=litellm_params, ) else: return asyncio.run( @@ -164,5 +167,6 @@ class AnthropicBatchesHandler: timeout=timeout, max_retries=max_retries, logging_obj=logging_obj, + litellm_params=litellm_params, ) ) diff --git a/litellm/llms/anthropic/batches/transformation.py b/litellm/llms/anthropic/batches/transformation.py index 7dee7513538..b95b734c761 100644 --- a/litellm/llms/anthropic/batches/transformation.py +++ b/litellm/llms/anthropic/batches/transformation.py @@ -13,6 +13,8 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.types.llms.openai import AllMessageValues, CreateBatchRequest from litellm.types.utils import LiteLLMBatch, LlmProviders, ModelResponse +from ..common_utils import merge_anthropic_beta_headers, without_caller_credential_headers + if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer @@ -69,24 +71,30 @@ class AnthropicBatchesConfig(BaseBatchesConfig): api_base: str | None = None, ) -> dict: """Validate and prepare environment-specific headers and parameters.""" - if api_base is None and isinstance(litellm_params, dict): - api_base = litellm_params.get("api_base") - auth_header: Final = self.anthropic_model_info.get_auth_header(api_key, api_base) + params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None + if api_base is None and params_mapping is not None: + api_base = params_mapping.get("api_base") + auth_header: Final = self.anthropic_model_info.get_auth_header( + api_key, api_base, litellm_params=params_mapping, allow_workload_identity=True + ) if auth_header is None: raise ValueError( "Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params" ) - _headers: Final = { + merged_beta: Final = merge_anthropic_beta_headers( + merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")), + "message-batches-2024-09-24", + ) + # The deployment's own credential is applied below, so a caller-supplied one must not + # ride along: without this a minted federation Bearer travels beside the caller's x-api-key. + return { + **without_caller_credential_headers(headers), "accept": "application/json", "anthropic-version": "2023-06-01", "content-type": "application/json", + **auth_header, + "anthropic-beta": merged_beta, } - _headers.update(auth_header) - # Add beta header for message batches - if "anthropic-beta" not in headers: - headers["anthropic-beta"] = "message-batches-2024-09-24" - headers.update(_headers) - return headers def get_complete_batch_url( self, diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 490912d42eb..fd67ebd5293 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -3,7 +3,7 @@ import re import time from collections.abc import Callable, Mapping, Sequence from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, NoReturn, cast +from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, cast import httpx from pydantic import BaseModel, ValidationError @@ -296,6 +296,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): to pass metadata to anthropic, it's {"user_id": "any-relevant-information"} """ + _workload_identity_eligible: ClassVar[bool] = True + max_tokens: int | None = None stop_sequences: list | None = None temperature: int | None = None diff --git a/litellm/llms/anthropic/common_utils.py b/litellm/llms/anthropic/common_utils.py index 65c2fccceeb..ee6db74fbe3 100644 --- a/litellm/llms/anthropic/common_utils.py +++ b/litellm/llms/anthropic/common_utils.py @@ -7,10 +7,11 @@ import re from collections.abc import Mapping, MutableMapping, Sequence from datetime import datetime, timezone from types import MappingProxyType -from typing import Any, Final, Literal, TypeVar +from typing import Any, ClassVar, Final, Literal, TypeVar +from urllib.parse import quote import httpx -from pydantic import BaseModel, ConfigDict, StrictBool, TypeAdapter, ValidationError +from pydantic import BaseModel, ConfigDict, Field, StrictBool, TypeAdapter, ValidationError import litellm from litellm.constants import ( @@ -28,8 +29,15 @@ from litellm.litellm_core_utils.prompt_templates.factory import ( THOUGHT_SIGNATURE_SEPARATOR, ) from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import message_field, parts_of +from litellm.llms.anthropic.wif import ( + aget_anthropic_wif_token, + anthropic_base_without_chat_suffix, + get_anthropic_wif_token, + warn_if_static_credential_shadows_federation, +) from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.proxy._types import SpecialHeaders from litellm.types.llms.anthropic import ( ANTHROPIC_HOSTED_TOOLS, ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER, @@ -222,6 +230,29 @@ def _strip_bedrock_id_suffixes(model: str) -> str: ) +_SERVER_OWNED_AUTH_HEADERS: Final = SpecialHeaders.litellm_credential_header_names() +_WIF_ELIGIBILITY_ATTR: Final = "_workload_identity_eligible" + + +def without_caller_credential_headers(headers: Mapping[str, str]) -> Mapping[str, str]: + """``headers`` minus every header that authenticates the caller to litellm. + + The deployment's own credential is applied on top of the result, so a caller-supplied + credential must not survive into the upstream request: without this a minted federation + Bearer travels beside the caller's own ``x-api-key``, and Anthropic sees two credentials. + """ + return MappingProxyType( + {name: value for name, value in headers.items() if name.lower() not in _SERVER_OWNED_AUTH_HEADERS} + ) + + +def config_allows_workload_identity(config: object) -> bool: + """A federation token is an Anthropic-org credential and its exchange POSTs the workload's OIDC + assertion to the deployment's own host, so eligibility is declared per class and read from that + class's own ``__dict__``: a subclass written for another provider inherits nothing.""" + return type(config).__dict__.get(_WIF_ELIGIBILITY_ATTR, False) is True + + def is_anthropic_oauth_key(value: str | None) -> bool: """Check if a value contains an Anthropic OAuth token (sk-ant-oat*).""" if value is None: @@ -240,12 +271,22 @@ def resolve_used_client_oauth_token(client_sent_oauth_token: object, custom_llm_ return client_sent_oauth_token and custom_llm_provider in ANTHROPIC_OAUTH_FORWARD_PROVIDERS -def _merge_beta_headers(existing: str | None, new_beta: str) -> str: - """Merge a new beta value into an existing comma-separated anthropic-beta header.""" - if not existing: - return new_beta - betas: Final = {b.strip() for b in existing.split(",") if b.strip()} - betas.add(new_beta) +def _beta_header_values(side: str | Sequence[str] | None) -> tuple[str, ...]: + if not side: + return () + if isinstance(side, str): + return (side,) + return tuple(entry for entry in side if isinstance(entry, str)) + + +def merge_anthropic_beta_headers(existing: str | Sequence[str] | None, new_beta: str | Sequence[str] | None) -> str: + """Merge anthropic-beta header values, deduplicated and sorted. + + Either side may arrive as a list rather than a comma-separated string: the Skills surface + accepted a list-valued header before it shared this helper, and callers still send one. + """ + joined: Final = ",".join(_beta_header_values(existing) + _beta_header_values(new_beta)) + betas: Final = frozenset(b.strip() for b in joined.split(",") if b.strip()) return ",".join(sorted(betas)) @@ -272,7 +313,9 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup ): headers.pop(name) headers["authorization"] = auth_header - headers["anthropic-beta"] = _merge_beta_headers(headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER) + headers["anthropic-beta"] = merge_anthropic_beta_headers( + headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER + ) headers["anthropic-dangerous-direct-browser-access"] = "true" return headers, api_key # Check api_key directly (standard chat/completion flow) @@ -280,7 +323,9 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup for name in tuple(header_name for header_name in headers if header_name.lower() == "x-api-key"): headers.pop(name) headers["authorization"] = f"Bearer {api_key}" - headers["anthropic-beta"] = _merge_beta_headers(headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER) + headers["anthropic-beta"] = merge_anthropic_beta_headers( + headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER + ) headers["anthropic-dangerous-direct-browser-access"] = "true" return headers, api_key @@ -316,7 +361,79 @@ class AnthropicError(BaseLLMException): super().__init__(status_code=status_code, message=message, headers=headers) +_MODEL_LIST_PAGE_CAP: Final = 20 + + +def _litellm_params_str(litellm_params: Mapping[str, object] | None, key: str) -> str | None: + value: Final = litellm_params.get(key) if litellm_params is not None else None + return value if isinstance(value, str) else None + + +class _AnthropicModelListEntry(BaseModel): + id: str + + +class _AnthropicModelsPage(BaseModel): + data: Sequence[_AnthropicModelListEntry] = Field(default_factory=tuple) + has_more: bool = False + last_id: str | None = None + + +def _sanitized_anthropic_error(response: httpx.Response, detail: str | None = None) -> str: + """A provider error detail built only from structured fields, never ``response.text`` + verbatim: the raw body is untrusted content the caller of ``/v1/models`` did not ask for + and should not have echoed back to it wholesale.""" + if detail is not None: + return f"HTTP {response.status_code}: {detail}" + try: + body: Final = response.json() + except ValueError: + return f"HTTP {response.status_code}" + error: Final = body.get("error") if isinstance(body, dict) else None + message: Final = error.get("message") if isinstance(error, dict) else None + return f"HTTP {response.status_code}: {message}" if isinstance(message, str) else f"HTTP {response.status_code}" + + +def _fetch_anthropic_models_page( + api_base: str, headers: Mapping[str, str], after_id: str | None +) -> _AnthropicModelsPage: + # after_id rides the URL because the client mutates the params mapping it is handed, + # which a read-only one cannot support + query: Final = f"?after_id={quote(after_id)}" if after_id else "" + response: Final = litellm.module_level_client.get( + url=f"{api_base}/v1/models{query}", + headers=headers, + follow_redirects=False, + ) + try: + response.raise_for_status() + except httpx.HTTPStatusError: + raise Exception(f"Failed to fetch models from Anthropic. {_sanitized_anthropic_error(response)}") from None + try: + return _AnthropicModelsPage.model_validate(response.json()) + except ValueError as e: + raise Exception( + f"Failed to fetch models from Anthropic. {_sanitized_anthropic_error(response, detail=str(e))}" + ) from None + + +def _fetch_anthropic_model_ids( + api_base: str, headers: Mapping[str, str], after_id: str | None, pages_left: int +) -> tuple[str, ...]: + collected: tuple[str, ...] = () # rebind-ok: accumulates one page of ids per iteration + cursor: str | None = after_id # rebind-ok: advances to each page's last_id + for _ in range(max(pages_left, 0)): + page = _fetch_anthropic_models_page(api_base, headers, cursor) + collected += tuple(entry.id for entry in page.data) + if not page.has_more or page.last_id is None: + return collected + cursor = page.last_id + raise Exception(f"Anthropic /v1/models did not terminate within {_MODEL_LIST_PAGE_CAP} pages.") + + class AnthropicModelInfo(BaseLLMModelInfo): + _workload_identity_eligible: ClassVar[bool] = True + def is_cache_control_set(self, messages: list[AllMessageValues]) -> bool: """ Return if {"cache_control": ..} in message content block @@ -940,7 +1057,9 @@ class AnthropicModelInfo(BaseLLMModelInfo): return list(set(betas).union(thinking_display_betas, tool_change_betas)) @staticmethod - def _make_api_key_auth_header(api_key: str, api_base: str | None, use_bearer_for_custom_base: bool = False) -> dict: + def _make_api_key_auth_header( + api_key: str, api_base: str | None, use_bearer_for_custom_base: bool = False + ) -> Mapping[str, str]: if use_bearer_for_custom_base and ( api_base and "api.anthropic.com" not in api_base and not api_key.startswith("sk-ant-") ): @@ -948,6 +1067,33 @@ class AnthropicModelInfo(BaseLLMModelInfo): return {"authorization": value} return {"x-api-key": api_key} + def _credential_headers( + self, + *, + api_key: str | None, + auth_token: str | None, + api_base: str | None, + use_bearer_for_custom_base: bool, + wif_minted: bool, + betas: set[str], # mutable-ok: the caller's beta accumulator, appended to by the oauth tier + ) -> Mapping[str, str]: + """The credential tier walk: a consumer OAuth token, then ANTHROPIC_AUTH_TOKEN, then an api key. + + A server-minted federation token takes the same Bearer shape as a consumer OAuth token but is + not browser-forwarded, so it does not get the direct-browser-access header. + """ + if api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX): + betas.add(ANTHROPIC_OAUTH_BETA_HEADER) + oauth_headers: Final = {"authorization": f"Bearer {api_key}"} + if wif_minted: + return oauth_headers + return {**oauth_headers, "anthropic-dangerous-direct-browser-access": "true"} + if auth_token and not api_key: + return {"authorization": f"Bearer {auth_token}"} + if api_key: + return self._make_api_key_auth_header(api_key, api_base, use_bearer_for_custom_base) + return {} + def get_anthropic_headers( self, api_key: str | None = None, @@ -972,6 +1118,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): is_mid_conversation_output_config_used: bool = False, is_thinking_display_updates_used: bool = False, is_mid_conversation_tool_change_used: bool = False, + wif_minted: bool = False, ) -> dict: betas: Final = set() # Anthropic no longer requires the prompt-caching beta header @@ -1010,20 +1157,21 @@ class AnthropicModelInfo(BaseLLMModelInfo): if is_mid_conversation_output_config_used: betas.add(ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER) - _is_oauth: Final = api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX) headers: Final = { "anthropic-version": anthropic_version or "2023-06-01", "accept": "application/json", "content-type": "application/json", } - if _is_oauth: - headers["authorization"] = f"Bearer {api_key}" - headers["anthropic-dangerous-direct-browser-access"] = "true" - betas.add(ANTHROPIC_OAUTH_BETA_HEADER) - elif auth_token and not api_key: - headers["authorization"] = f"Bearer {auth_token}" - elif api_key: - headers.update(self._make_api_key_auth_header(api_key, api_base, use_bearer_for_custom_base)) + headers.update( + self._credential_headers( + api_key=api_key, + auth_token=auth_token, + api_base=api_base, + use_bearer_for_custom_base=use_bearer_for_custom_base, + wif_minted=wif_minted, + betas=betas, + ) + ) if user_anthropic_beta_headers is not None: betas.update(user_anthropic_beta_headers) @@ -1055,10 +1203,11 @@ class AnthropicModelInfo(BaseLLMModelInfo): api_key: str | None = None, api_base: str | None = None, ) -> dict: - if api_base is None and isinstance(litellm_params, dict): - api_base = litellm_params.get("api_base") + params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None + if api_base is None and params_mapping is not None: + api_base = params_mapping.get("api_base") use_bearer_for_custom_base: Final[bool] = bool( - isinstance(litellm_params, dict) and litellm_params.get("use_bearer_for_custom_base", False) + params_mapping is not None and params_mapping.get("use_bearer_for_custom_base", False) ) # Check for Anthropic OAuth token in headers headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key) @@ -1067,9 +1216,25 @@ class AnthropicModelInfo(BaseLLMModelInfo): auth_token: str | None = None if api_key is None: auth_token = AnthropicModelInfo.get_auth_token() - if api_key is None and auth_token is None: + if (api_key is not None or auth_token is not None) and config_allows_workload_identity(self): + warn_if_static_credential_shadows_federation(params_mapping, model) + wif_token: Final = ( + get_anthropic_wif_token(params_mapping, api_base, model) + if api_key is None and auth_token is None and config_allows_workload_identity(self) + else None + ) + wif_minted: Final = wif_token is not None + resolved_api_key: Final = wif_token if wif_token is not None else api_key + if resolved_api_key is None and auth_token is None: raise litellm.AuthenticationError( - message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` in your environment vars", + message=( + "Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the " + "environment variables or via params. Please set `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` " + "in your environment vars, or configure workload identity federation via " + "`ANTHROPIC_FEDERATION_RULE_ID`, `ANTHROPIC_ORGANIZATION_ID`, " + "`ANTHROPIC_SERVICE_ACCOUNT_ID` and " + "`ANTHROPIC_IDENTITY_TOKEN_FILE` (or `ANTHROPIC_IDENTITY_TOKEN`)" + ), llm_provider="anthropic", model=model, ) @@ -1095,7 +1260,7 @@ class AnthropicModelInfo(BaseLLMModelInfo): computer_tool_used=computer_tool_used, prompt_caching_set=prompt_caching_set, pdf_used=pdf_used, - api_key=api_key, + api_key=resolved_api_key, auth_token=auth_token, file_id_used=file_id_used, is_mid_conversation_output_config_used=is_mid_conversation_output_config_used, @@ -1113,11 +1278,12 @@ class AnthropicModelInfo(BaseLLMModelInfo): container_with_skills_used=container_with_skills_used, api_base=api_base, use_bearer_for_custom_base=use_bearer_for_custom_base, + wif_minted=wif_minted, ) - headers = {**headers, **anthropic_headers} + caller_headers: Final = without_caller_credential_headers(headers) if wif_minted else headers - return headers + return {**caller_headers, **anthropic_headers} @staticmethod def get_api_base(api_base: str | None = None) -> str | None: @@ -1132,9 +1298,13 @@ class AnthropicModelInfo(BaseLLMModelInfo): @staticmethod def get_api_key(api_key: str | None = None) -> str | None: - from litellm.secret_managers.main import get_secret_str + """An empty or whitespace-only key counts as unset: it can never authenticate anything, and + treating it as set would silently outrank workload identity federation.""" + from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str - return api_key or get_secret_str("ANTHROPIC_API_KEY") + return normalize_nonempty_secret_str(api_key) or normalize_nonempty_secret_str( + get_secret_str("ANTHROPIC_API_KEY") + ) @staticmethod def get_auth_token(auth_token: str | None = None) -> str | None: @@ -1143,61 +1313,130 @@ class AnthropicModelInfo(BaseLLMModelInfo): Unlike api_key (which uses X-Api-Key header), auth_token uses Authorization: Bearer header, matching the official Anthropic SDK behavior. """ - from litellm.secret_managers.main import get_secret_str + from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str - return auth_token or get_secret_str("ANTHROPIC_AUTH_TOKEN") + return normalize_nonempty_secret_str(auth_token) or normalize_nonempty_secret_str( + get_secret_str("ANTHROPIC_AUTH_TOKEN") + ) @staticmethod def get_auth_header( api_key: str | None = None, api_base: str | None = None, use_bearer_for_custom_base: bool = False, - ) -> dict | None: + litellm_params: Mapping[str, object] | None = None, + allow_workload_identity: bool = False, + ) -> Mapping[str, str] | None: """Resolve Anthropic credentials and return the appropriate auth header dict. Checks ANTHROPIC_API_KEY first (-> x-api-key or Bearer depending on - use_bearer_for_custom_base), then ANTHROPIC_AUTH_TOKEN (-> Authorization: Bearer). - Returns None if neither is available. + use_bearer_for_custom_base), then ANTHROPIC_AUTH_TOKEN (-> Authorization: Bearer), + then workload identity federation (-> Authorization: Bearer with a minted + sk-ant-oat01 token, honoring anthropic_* litellm_params when provided). Every + Bearer built from an sk-ant-oat token carries the mandatory oauth anthropic-beta. + Returns None if no credential source is available. """ + static_header: Final = AnthropicModelInfo._static_auth_header(api_key, api_base, use_bearer_for_custom_base) + if static_header is not None: + return static_header + if not allow_workload_identity: + return None + wif_token: Final = get_anthropic_wif_token(litellm_params, api_base, "") + if wif_token is not None: + return AnthropicModelInfo._oauth_bearer_header(wif_token) + return None + + @staticmethod + async def aget_auth_header( + api_key: str | None = None, + api_base: str | None = None, + use_bearer_for_custom_base: bool = False, + litellm_params: Mapping[str, object] | None = None, + allow_workload_identity: bool = False, + ) -> Mapping[str, str] | None: + """Async counterpart of get_auth_header: the WIF tier can block on a token + exchange POST, so async callers await it off the event loop.""" + static_header: Final = AnthropicModelInfo._static_auth_header(api_key, api_base, use_bearer_for_custom_base) + if static_header is not None: + return static_header + if not allow_workload_identity: + return None + wif_token: Final = await aget_anthropic_wif_token(litellm_params, api_base, "") + if wif_token is not None: + return AnthropicModelInfo._oauth_bearer_header(wif_token) + return None + + @staticmethod + def _static_auth_header( + api_key: str | None, + api_base: str | None, + use_bearer_for_custom_base: bool, + ) -> Mapping[str, str] | None: resolved_key: Final = AnthropicModelInfo.get_api_key(api_key) if resolved_key is not None: if is_anthropic_oauth_key(resolved_key): - return {"authorization": f"Bearer {resolved_key}"} + return AnthropicModelInfo._oauth_bearer_header(resolved_key) return AnthropicModelInfo._make_api_key_auth_header(resolved_key, api_base, use_bearer_for_custom_base) auth_token: Final = AnthropicModelInfo.get_auth_token() if auth_token is not None: return {"authorization": f"Bearer {auth_token}"} return None + @staticmethod + def _oauth_bearer_header(token: str) -> Mapping[str, str]: + return {"authorization": f"Bearer {token}", "anthropic-beta": ANTHROPIC_OAUTH_BETA_HEADER} + @staticmethod def get_base_model(model: str | None = None) -> str | None: return model.replace("anthropic/", "") if model else None def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]: - api_base = AnthropicModelInfo.get_api_base(api_base) - auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base) - if api_base is None or auth_header is None: - raise ValueError( - "ANTHROPIC_API_BASE/ANTHROPIC_BASE_URL or ANTHROPIC_API_KEY/ANTHROPIC_AUTH_TOKEN is not set. Please set the environment variable, to query Anthropic's `/models` endpoint." - ) - headers: Final = {"anthropic-version": "2023-06-01"} - headers.update(auth_header) - response: Final = litellm.module_level_client.get( - url=f"{api_base}/v1/models", - headers=headers, + return self._list_models(api_key=api_key, api_base=api_base, litellm_params=None) + + def discover_models( + self, litellm_params: Mapping[str, object] | None = None + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + """Live discovery for a configured deployment: unlike ``get_models``, this threads the + full ``litellm_params`` into ``get_auth_header`` so a workload-identity-federation source + configured on the deployment (rather than the environment) is honored, gated the same way + every other Anthropic auth surface is via ``config_allows_workload_identity``.""" + return self._list_models( + api_key=_litellm_params_str(litellm_params, "api_key"), + api_base=_litellm_params_str(litellm_params, "api_base"), + litellm_params=litellm_params, ) - try: - response.raise_for_status() - except httpx.HTTPStatusError: - raise Exception( - f"Failed to fetch models from Anthropic. Status code: {response.status_code}, Response: {response.text}" + def _list_models( + self, + *, + api_key: str | None, + api_base: str | None, + litellm_params: Mapping[str, object] | None, + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + resolved_api_base: Final = AnthropicModelInfo.get_api_base(api_base) + auth_header: Final = AnthropicModelInfo.get_auth_header( + api_key, + resolved_api_base, + litellm_params=litellm_params, + allow_workload_identity=config_allows_workload_identity(self), + ) + if resolved_api_base is None or auth_header is None: + raise ValueError( + "ANTHROPIC_API_BASE/ANTHROPIC_BASE_URL or ANTHROPIC_API_KEY/ANTHROPIC_AUTH_TOKEN (or workload " + "identity federation via ANTHROPIC_FEDERATION_RULE_ID/ANTHROPIC_ORGANIZATION_ID/" + "ANTHROPIC_IDENTITY_TOKEN_FILE) is not set. Please set the environment variable, to query " + "Anthropic's `/models` endpoint." ) - - models: Final[Sequence[Mapping[str, str]]] = response.json()["data"] - - litellm_model_names: Final = ["anthropic/" + model["id"] for model in models] - return litellm_model_names + headers: Final = MappingProxyType({"anthropic-version": "2023-06-01", **auth_header}) + # /v1/models is appended below, so a base the operator already wrote as .../v1 or + # .../v1/messages would otherwise be asked for /v1/v1/models. + model_ids: Final = _fetch_anthropic_model_ids( + anthropic_base_without_chat_suffix(resolved_api_base), + headers, + after_id=None, + pages_left=_MODEL_LIST_PAGE_CAP, + ) + return ["anthropic/" + model_id for model_id in model_ids] def get_token_counter(self) -> BaseTokenCounter | None: """ diff --git a/litellm/llms/anthropic/count_tokens/handler.py b/litellm/llms/anthropic/count_tokens/handler.py index dd2135f4918..b4f107cdef4 100644 --- a/litellm/llms/anthropic/count_tokens/handler.py +++ b/litellm/llms/anthropic/count_tokens/handler.py @@ -32,7 +32,7 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): self, model: str, messages: list[dict[str, JsonValue]], - api_key: str, + auth_header: Mapping[str, str], api_base: str | None = None, timeout: float | httpx.Timeout | None = None, tools: list[dict[str, JsonValue]] | None = None, @@ -45,8 +45,8 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): Args: model: The model identifier (e.g., "claude-3-5-sonnet-20241022") messages: The messages to count tokens for - api_key: The Anthropic API key - api_base: Optional custom API base URL + auth_header: The resolved Anthropic auth header (``AnthropicModelInfo.get_auth_header``) + api_base: Optional deployment api_base the count-tokens path is appended to timeout: Optional timeout for the request (defaults to litellm.request_timeout) Returns: @@ -73,12 +73,12 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig): verbose_logger.debug("Transformed request: %s", request_body) # Get endpoint URL - endpoint_url: Final = api_base or self.get_anthropic_count_tokens_endpoint() + endpoint_url: Final = self.get_anthropic_count_tokens_endpoint(api_base) verbose_logger.debug("Making request to: %s", endpoint_url) # Get required headers - headers: Final = self.get_required_headers(api_key) + headers: Final = self.get_count_tokens_headers(auth_header) # Use LiteLLM's async httpx client async_client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.ANTHROPIC) diff --git a/litellm/llms/anthropic/count_tokens/token_counter.py b/litellm/llms/anthropic/count_tokens/token_counter.py index 8e8d10c961b..920916726f2 100644 --- a/litellm/llms/anthropic/count_tokens/token_counter.py +++ b/litellm/llms/anthropic/count_tokens/token_counter.py @@ -2,10 +2,10 @@ Anthropic Token Counter implementation using the CountTokens API. """ -import os from typing import Any, Final from litellm._logging import verbose_logger +from litellm.exceptions import AuthenticationError from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler from litellm.llms.base_llm.base_utils import BaseTokenCounter from litellm.types.utils import LlmProviders, TokenCountResponse @@ -46,28 +46,31 @@ class AnthropicTokenCounter(BaseTokenCounter): Returns: TokenCountResponse with token count, or None if counting fails """ - from litellm.llms.anthropic.common_utils import AnthropicError + from litellm.llms.anthropic.common_utils import AnthropicError, AnthropicModelInfo if not messages: return None deployment = deployment or {} litellm_params: Final = deployment.get("litellm_params", {}) - - # Get Anthropic API key from deployment config or environment - api_key = litellm_params.get("api_key") - if not api_key: - api_key = os.getenv("ANTHROPIC_API_KEY") - - if not api_key: - verbose_logger.warning("No Anthropic API key found for token counting") - return None + api_base: Final = litellm_params.get("api_base") try: + auth_header: Final = await AnthropicModelInfo.aget_auth_header( + api_key=litellm_params.get("api_key"), + api_base=api_base, + litellm_params=litellm_params, + allow_workload_identity=True, + ) + if auth_header is None: + verbose_logger.warning("No Anthropic credential found for token counting") + return None + result: Final = await anthropic_count_tokens_handler.handle_count_tokens_request( model=model_to_use, messages=messages, - api_key=api_key, + auth_header=auth_header, + api_base=api_base, tools=tools, system=system, ) @@ -80,8 +83,8 @@ class AnthropicTokenCounter(BaseTokenCounter): tokenizer_type="anthropic_api", original_response=result, ) - except AnthropicError as e: - verbose_logger.warning("Anthropic CountTokens API error: status=%s, message=%s", e.status_code, e.message) + except (AnthropicError, AuthenticationError) as e: + verbose_logger.warning("Anthropic CountTokens error: status=%s, message=%s", e.status_code, e.message) return TokenCountResponse( total_tokens=0, request_model=request_model, diff --git a/litellm/llms/anthropic/count_tokens/transformation.py b/litellm/llms/anthropic/count_tokens/transformation.py index fb12747cec0..1e97d56913a 100644 --- a/litellm/llms/anthropic/count_tokens/transformation.py +++ b/litellm/llms/anthropic/count_tokens/transformation.py @@ -11,6 +11,8 @@ from typing import Final from pydantic import JsonValue, TypeAdapter from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION +from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers +from litellm.llms.anthropic.wif import resolve_anthropic_base _COUNT_REQUEST: Final = TypeAdapter(dict[str, JsonValue]) COUNT_TOKEN_OPTION_NAMES: Final = ("thinking", "tool_choice", "output_config") @@ -26,14 +28,21 @@ class AnthropicCountTokensConfig: - Response: {"input_tokens": } """ - def get_anthropic_count_tokens_endpoint(self) -> str: + def get_anthropic_count_tokens_endpoint(self, api_base: str | None = None) -> str: """ Get the Anthropic CountTokens API endpoint. + Args: + api_base: The deployment's api_base, which names the chat surface (a host, or a + base already carrying ``/v1`` or ``/v1/messages``); the count-tokens path is + appended to it, so it is never the full count-tokens URL. Unset or empty falls + back to ``ANTHROPIC_API_BASE`` / ``ANTHROPIC_BASE_URL`` and then Anthropic's + host, the same resolution chat and the federated exchange use + Returns: The endpoint URL for the CountTokens API """ - return "https://api.anthropic.com/v1/messages/count_tokens" + return resolve_anthropic_base(api_base) + "/v1/messages/count_tokens" def transform_request_to_count_tokens( self, @@ -64,28 +73,19 @@ class AnthropicCountTokensConfig: ) ) - def get_required_headers(self, api_key: str) -> dict[str, str]: - """ - Get the required headers for the CountTokens API. - - Args: - api_key: The Anthropic API key - - Returns: - Dictionary of required headers - """ - from litellm.llms.anthropic.common_utils import ( - optionally_handle_anthropic_oauth, - ) - - headers: dict[str, str] = { + def get_count_tokens_headers(self, auth_header: Mapping[str, str]) -> dict[str, str]: + """The count-tokens headers around a resolved Anthropic auth header + (``AnthropicModelInfo.get_auth_header``): x-api-key for a static key, an Authorization + bearer for ``ANTHROPIC_AUTH_TOKEN`` and for sk-ant-oat tokens, whose mandatory oauth beta + merges with the token-counting beta instead of replacing it.""" + return { "Content-Type": "application/json", - "x-api-key": api_key, "anthropic-version": "2023-06-01", - "anthropic-beta": ANTHROPIC_TOKEN_COUNTING_BETA_VERSION, + **auth_header, + "anthropic-beta": merge_anthropic_beta_headers( + auth_header.get("anthropic-beta"), ANTHROPIC_TOKEN_COUNTING_BETA_VERSION + ), } - headers, _ = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key) - return headers def validate_request( self, diff --git a/litellm/llms/anthropic/files/handler.py b/litellm/llms/anthropic/files/handler.py index e4c75a704ec..f1f9c204b58 100644 --- a/litellm/llms/anthropic/files/handler.py +++ b/litellm/llms/anthropic/files/handler.py @@ -1,7 +1,7 @@ import asyncio import json import time -from collections.abc import Coroutine +from collections.abc import Coroutine, Mapping from typing import Final import httpx @@ -43,6 +43,7 @@ class AnthropicFilesHandler: api_key: str | None = None, timeout: float | httpx.Timeout = 600.0, max_retries: int | None = None, + litellm_params: Mapping[str, object] | None = None, ) -> HttpxBinaryResponseContent: """ Async: Retrieve file content from Anthropic. @@ -56,6 +57,7 @@ class AnthropicFilesHandler: api_key: Anthropic API key timeout: Request timeout max_retries: Max retry attempts (unused for now) + litellm_params: Deployment params, so a named credential's federation settings reach the mint Returns: HttpxBinaryResponseContent: Binary content wrapped in compatible response format @@ -73,7 +75,9 @@ class AnthropicFilesHandler: # Get Anthropic API credentials api_base = self.anthropic_model_info.get_api_base(api_base) - auth_header: Final = self.anthropic_model_info.get_auth_header(api_key, api_base) + auth_header: Final = await self.anthropic_model_info.aget_auth_header( + api_key, api_base, litellm_params=litellm_params, allow_workload_identity=True + ) if auth_header is None: raise ValueError("Missing Anthropic API Key") @@ -116,6 +120,7 @@ class AnthropicFilesHandler: api_key: str | None = None, timeout: float | httpx.Timeout = 600.0, max_retries: int | None = None, + litellm_params: Mapping[str, object] | None = None, ) -> HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent]: """ Retrieve file content from Anthropic. @@ -130,6 +135,7 @@ class AnthropicFilesHandler: api_key: Anthropic API key timeout: Request timeout max_retries: Max retry attempts (unused for now) + litellm_params: Deployment params, so a named credential's federation settings reach the mint Returns: HttpxBinaryResponseContent or Coroutine: Binary content wrapped in compatible response format @@ -139,7 +145,9 @@ class AnthropicFilesHandler: file_content_request=file_content_request, api_base=api_base, api_key=api_key, + timeout=timeout, max_retries=max_retries, + litellm_params=litellm_params, ) else: return asyncio.run( @@ -149,6 +157,7 @@ class AnthropicFilesHandler: api_key=api_key, timeout=timeout, max_retries=max_retries, + litellm_params=litellm_params, ) ) diff --git a/litellm/llms/anthropic/files/transformation.py b/litellm/llms/anthropic/files/transformation.py index 7b5ab78af8d..04d057ed3e9 100644 --- a/litellm/llms/anthropic/files/transformation.py +++ b/litellm/llms/anthropic/files/transformation.py @@ -14,6 +14,7 @@ Anthropic Files API endpoints: import calendar import time +from collections.abc import Mapping from typing import Final, cast import httpx @@ -35,7 +36,12 @@ from litellm.types.llms.openai import ( ) from litellm.types.utils import LlmProviders -from ..common_utils import AnthropicError, AnthropicModelInfo +from ..common_utils import ( + AnthropicError, + AnthropicModelInfo, + merge_anthropic_beta_headers, + without_caller_credential_headers, +) ANTHROPIC_FILES_API_BASE: Final = "https://api.anthropic.com" ANTHROPIC_FILES_BETA_HEADER: Final = "files-api-2025-04-14" @@ -94,21 +100,55 @@ class AnthropicFilesConfig(BaseFilesConfig): api_key: str | None = None, api_base: str | None = None, ) -> dict: - if api_base is None and isinstance(litellm_params, dict): - api_base = litellm_params.get("api_base") - auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base) + params_mapping, resolved_api_base = self._resolve_params(litellm_params, api_base) + auth_header: Final = AnthropicModelInfo.get_auth_header( + api_key, resolved_api_base, litellm_params=params_mapping, allow_workload_identity=True + ) + return self._finalize_headers(headers, auth_header) + + async def avalidate_environment( + self, + headers: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + model: str, + messages: list, # mutable-ok: mirrors the sync validate_environment contract this overrides + optional_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + litellm_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: # mutable-ok: mirrors the sync validate_environment contract this overrides + """Async counterpart of validate_environment: the WIF tier can block on a token + exchange POST, so async callers await it off the event loop.""" + params_mapping, resolved_api_base = self._resolve_params(litellm_params, api_base) + auth_header: Final = await AnthropicModelInfo.aget_auth_header( + api_key, resolved_api_base, litellm_params=params_mapping, allow_workload_identity=True + ) + return self._finalize_headers(headers, auth_header) + + @staticmethod + def _resolve_params( + litellm_params: dict, api_base: str | None + ) -> tuple[dict | None, str | None]: # mutable-ok: mirrors the sync validate_environment contract this overrides + params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None + if api_base is None and params_mapping is not None: + api_base = params_mapping.get("api_base") + return params_mapping, api_base + + @staticmethod + def _finalize_headers(headers: dict, auth_header: Mapping[str, str] | None) -> dict: # mutable-ok: out-param if auth_header is None: raise ValueError( "Anthropic API key is required. Set ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN environment variable or pass api_key parameter." ) - headers.update( - { - **auth_header, - "anthropic-version": "2023-06-01", - "anthropic-beta": ANTHROPIC_FILES_BETA_HEADER, - } + merged_beta: Final = merge_anthropic_beta_headers( + merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")), + ANTHROPIC_FILES_BETA_HEADER, ) - return headers + return { + **without_caller_credential_headers(headers), + **auth_header, + "anthropic-version": "2023-06-01", + "anthropic-beta": merged_beta, + } def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]: return ["purpose"] diff --git a/litellm/llms/anthropic/pass_through/messages/transformation.py b/litellm/llms/anthropic/pass_through/messages/transformation.py index 2fbb51ec949..bc61b133f4e 100644 --- a/litellm/llms/anthropic/pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/pass_through/messages/transformation.py @@ -1,5 +1,5 @@ from collections.abc import AsyncIterator, Mapping, Sequence -from typing import Any, Final +from typing import Any, ClassVar, Final import httpx @@ -25,6 +25,7 @@ from litellm.types.router import GenericLiteLLMParams from ...common_utils import ( AnthropicError, AnthropicModelInfo, + merge_anthropic_beta_headers, optionally_handle_anthropic_oauth, requires_native_compaction_beta, strip_advisor_blocks_from_messages, @@ -38,6 +39,17 @@ from .mid_conversation_system import ( DEFAULT_ANTHROPIC_API_VERSION: Final = "2023-06-01" +_CALLER_CREDENTIAL_HEADERS: Final = frozenset({"x-api-key", "authorization"}) + + +def _carries_caller_credential(headers: Mapping[str, str]) -> bool: + """Whether the caller sent their own Anthropic credential, in which case this passthrough + honors it and never mints. Matched case-insensitively: an SDK caller passing ``X-Api-Key`` + through extra_headers would otherwise slip the check and end up sending their key beside a + minted federation Bearer.""" + return any(name.lower() in _CALLER_CREDENTIAL_HEADERS for name in headers) + + DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING: Final = ( "Dropping adaptive `thinking`/`output_config.effort` for model=%s: the model " "does not support extended thinking, or max_tokens is too small to fit the " @@ -55,6 +67,8 @@ def _messages_carry_output_config(messages: Sequence[object]) -> bool: class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): + _workload_identity_eligible: ClassVar[bool] = True + @property def custom_llm_provider(self) -> str | None: return "anthropic" @@ -256,33 +270,109 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): # Check for Anthropic OAuth token in Authorization header headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key) - header_names: Final = frozenset(name.lower() for name in headers) - if "x-api-key" not in header_names and "authorization" not in header_names: - auth_header: Final = AnthropicModelInfo.get_auth_header(api_key) - if auth_header is None: - raise AuthenticationError( - message=( - "Missing Anthropic API Key - A call is being made to anthropic but no key is set " - "either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` " - "or `ANTHROPIC_AUTH_TOKEN` in your environment vars" + if not _carries_caller_credential(headers): + self._apply_env_auth_header( + headers, + self._require_auth_header( + AnthropicModelInfo.get_auth_header( + api_key, + api_base=api_base, + litellm_params=litellm_params, + allow_workload_identity=self._allows_workload_identity, ), - llm_provider=self._resolved_provider, model=model, - ) - headers.update(auth_header) + ), + ) + return self._finalize_messages_headers(headers, optional_params, messages), api_base + + async def avalidate_anthropic_messages_environment( + self, + headers: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + model: str, + messages: list[Any], # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + optional_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + litellm_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + if type(self).validate_anthropic_messages_environment is not ( + AnthropicMessagesConfig.validate_anthropic_messages_environment + ): + # a subclass sync override must keep winning on the async path + return self.validate_anthropic_messages_environment( + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=api_key, + api_base=api_base, + ) + oauth_headers, oauth_api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key) + + if not _carries_caller_credential(oauth_headers): + self._apply_env_auth_header( + oauth_headers, + self._require_auth_header( + await AnthropicModelInfo.aget_auth_header( + oauth_api_key, + api_base=api_base, + litellm_params=litellm_params, + allow_workload_identity=self._allows_workload_identity, + ), + model=model, + ), + ) + return self._finalize_messages_headers(oauth_headers, optional_params, messages), api_base + + def _require_auth_header(self, auth_header: Mapping[str, str] | None, model: str) -> Mapping[str, str]: + if auth_header is None: + raise AuthenticationError( + message=( + "Missing Anthropic API Key - A call is being made to anthropic but no key is set " + "either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` " + "or `ANTHROPIC_AUTH_TOKEN` in your environment vars" + ), + llm_provider=self._resolved_provider, + model=model, + ) + return auth_header + + @staticmethod + def _apply_env_auth_header(headers: dict, auth_header: Mapping[str, str] | None) -> None: # mutable-ok: out-param + if auth_header is None: + return + merged_beta: Final = merge_anthropic_beta_headers( + headers.get("anthropic-beta"), auth_header.get("anthropic-beta") + ) + headers.update(auth_header) + if merged_beta: + headers["anthropic-beta"] = merged_beta + + @property + def _allows_workload_identity(self) -> bool: + """Subclasses reuse this validate step for their own /v1/messages-compatible providers, so + eligibility is declared per class and never inherited.""" + from litellm.llms.anthropic.common_utils import config_allows_workload_identity + + return config_allows_workload_identity(self) + + def _finalize_messages_headers( + self, + headers: dict, # mutable-ok: out-param + optional_params: dict, # mutable-ok: out-param + messages: list[Any], # mutable-ok: mirrors the validate_anthropic_messages_environment contract + ) -> dict: # mutable-ok: out-param if "anthropic-version" not in headers: headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION if "content-type" not in headers: headers["content-type"] = "application/json" - - headers = self._update_headers_with_anthropic_beta( + return self._update_headers_with_anthropic_beta( headers=headers, optional_params=optional_params, messages=messages, ) - return headers, api_base - @staticmethod def _translate_reasoning_effort_to_anthropic( model: str, optional_params: dict, max_tokens: int | None, custom_llm_provider: str diff --git a/litellm/llms/anthropic/prompt_cache_prediction.py b/litellm/llms/anthropic/prompt_cache_prediction.py index 6528c7ac726..ff304e845f7 100644 --- a/litellm/llms/anthropic/prompt_cache_prediction.py +++ b/litellm/llms/anthropic/prompt_cache_prediction.py @@ -524,17 +524,19 @@ async def count_prompt_tokens( body: Mapping[str, JsonValue], api_base: str | None = None, ) -> int | None: + auth_header: Final = AnthropicModelInfo.get_auth_header(api_key=api_key, api_base=api_base) + if auth_header is None: + return None try: native: Final = _CountBody.model_validate(body) - count_url: Final = _messages_url(model, api_key, api_base) + "/count_tokens" result: Final = _CountResult.model_validate( await _counter.handle_count_tokens_request( model=model, messages=_count_objects(native.messages), tools=_count_objects(native.tools) if native.tools is not None else None, system=_JSON_OBJECT.validate_python(MappingProxyType({"system": native.system}))["system"], - api_key=api_key, - api_base=count_url, + auth_header=auth_header, + api_base=api_base, optional_params=_JSON_OBJECT.validate_python( MappingProxyType({key: body[key] for key in COUNT_TOKEN_OPTION_NAMES if key in body}) ), diff --git a/litellm/llms/anthropic/skills/transformation.py b/litellm/llms/anthropic/skills/transformation.py index 448e2dc2584..6d419a56889 100644 --- a/litellm/llms/anthropic/skills/transformation.py +++ b/litellm/llms/anthropic/skills/transformation.py @@ -2,6 +2,7 @@ Anthropic Skills API configuration and transformations """ +from types import MappingProxyType from typing import Final import httpx @@ -35,40 +36,35 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig): def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict: """Add Anthropic-specific headers""" - from litellm.llms.anthropic.common_utils import AnthropicModelInfo + from litellm.constants import ANTHROPIC_SKILLS_API_BETA_VERSION + from litellm.llms.anthropic.common_utils import ( + AnthropicModelInfo, + merge_anthropic_beta_headers, + without_caller_credential_headers, + ) - # Get API key from litellm_params if available - api_key = None - api_base = None - if litellm_params is not None: - api_key = litellm_params.api_key - api_base = litellm_params.api_base - - auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base) + auth_header: Final = AnthropicModelInfo.get_auth_header( + api_key=litellm_params.api_key if litellm_params is not None else None, + api_base=litellm_params.api_base if litellm_params is not None else None, + litellm_params=MappingProxyType(dict(litellm_params)) if litellm_params is not None else None, + allow_workload_identity=True, + ) if auth_header is None: raise ValueError("ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN is required for Skills API") - headers.update(auth_header) - headers["anthropic-version"] = "2023-06-01" - - # Add beta header for skills API - from litellm.constants import ANTHROPIC_SKILLS_API_BETA_VERSION - - if "anthropic-beta" not in headers: - headers["anthropic-beta"] = ANTHROPIC_SKILLS_API_BETA_VERSION - elif isinstance(headers["anthropic-beta"], list): - if ANTHROPIC_SKILLS_API_BETA_VERSION not in headers["anthropic-beta"]: - headers["anthropic-beta"].append(ANTHROPIC_SKILLS_API_BETA_VERSION) - elif isinstance(headers["anthropic-beta"], str): - if ANTHROPIC_SKILLS_API_BETA_VERSION not in headers["anthropic-beta"]: - headers["anthropic-beta"] = [ - headers["anthropic-beta"], - ANTHROPIC_SKILLS_API_BETA_VERSION, - ] - - headers["content-type"] = "application/json" - - return headers + merged_beta: Final = merge_anthropic_beta_headers( + merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")), + ANTHROPIC_SKILLS_API_BETA_VERSION, + ) + # The deployment's own credential is applied here, so a caller-supplied one must not ride + # along upstream beside a minted federation Bearer. + return { + **without_caller_credential_headers(headers), + **auth_header, + "anthropic-version": "2023-06-01", + "anthropic-beta": merged_beta, + "content-type": "application/json", + } def get_complete_url( self, diff --git a/litellm/llms/anthropic/wif.py b/litellm/llms/anthropic/wif.py new file mode 100644 index 00000000000..b096a11ace4 --- /dev/null +++ b/litellm/llms/anthropic/wif.py @@ -0,0 +1,604 @@ +"""Anthropic workload identity federation: exchanges an external OIDC identity +token for a short-lived ``sk-ant-oat01`` token via the shared RFC 7523 engine.""" + +import os +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from functools import lru_cache +from itertools import chain +from types import MappingProxyType +from typing import Final, NoReturn, TypeVar +from urllib.parse import urlsplit, urlunsplit + +from pydantic import BaseModel, ConfigDict, ValidationError +from typing_extensions import assert_never + +import litellm +from litellm._logging import verbose_logger +from litellm.llms.base_llm.auth.client_credentials import keycloak_assertion_source +from litellm.llms.base_llm.auth.identity_source import ( + AnthropicIdentitySourceKind, + InternalIssuerSource, + KeycloakSource, + identity_source_ref, +) +from litellm.llms.base_llm.auth.internal_issuer import ( + internal_issuer_assertion_source, + internal_issuer_jwks_document, +) +from litellm.llms.base_llm.auth.token_exchange import ( + JwtBearerTokenExchangeEngine, + default_token_exchange_engine, +) +from litellm.llms.base_llm.auth.types import ( + AssertionSourceError, + ExchangeError, + ExchangeResult, + InsecureTokenUrl, + MalformedTokenResponse, + MintedToken, + TokenEndpointError, + TokenExchangeSpec, + TokenTransportError, +) +from litellm.types.llms.anthropic import ANTHROPIC_TOKEN_EXCHANGE_PATH + +_JWT_BEARER_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer" +_DEFAULT_API_BASE: Final = "https://api.anthropic.com" +_INLINE_ENV_VAR: Final = "ANTHROPIC_IDENTITY_TOKEN" +_DISABLE_WIF_PARAM: Final = "anthropic_disable_workload_identity_federation" +_ACCEPTED_REF_PREFIX: Final = "oidc/" +_SHADOWED_DEPLOYMENT_WARNING_CAP: Final = 512 +_CHAT_BASE_SUFFIXES: Final = ("/v1/messages", "/v1") +# Hosts a federated exchange may talk to. api_base decides where the workload's assertion is sent +# AND where the minted org-scoped token is presented, so anyone able to write api_base on a +# federated deployment could otherwise redirect both. Gating each write path does not terminate: +# a deployment, a referenced credential and a future endpoint all reach the same value. This is the +# one place a federated exchange is built, so the trust decision is enforced here instead, and the +# allowlist is server-owned -- read from the environment, never from a model or credential API. +_TRUSTED_EXCHANGE_HOSTS_ENV: Final = "LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS" +_SCHEME_DEFAULT_PORTS: Final[Mapping[str, int]] = MappingProxyType({"http": 80, "https": 443}) +_DEFAULT_TRUSTED_EXCHANGE_HOST: Final = "api.anthropic.com" +_REJECTED_REF_PREFIX: Final = "oidc/env_path/" +_IDENTITY_SOURCE_PARAM: Final = "anthropic_identity_source" +_IDENTITY_SOURCE_ENV: Final = "ANTHROPIC_IDENTITY_SOURCE" +_IDENTITY_TOKEN_FILE_PARAM: Final = "anthropic_identity_token_file" +_IDENTITY_TOKEN_PARAM: Final = "anthropic_identity_token" + +# litellm_params key -> InternalIssuerSource/KeycloakSource field name. Every key here must +# also be listed in ANTHROPIC_WIF_KWARGS_KEYS (types/workload_identity.py), which is what makes it +# request-banned and cleared on a client-redirected api_base -- see types/utils.py's +# anthropic_wif_litellm_params, derived from that same set. +_INTERNAL_ISSUER_FIELD_MAP: Final[Mapping[str, str]] = MappingProxyType( + { + "anthropic_issuer_url": "issuer_url", + "anthropic_issuer_subject": "subject", + "anthropic_issuer_audience": "audience", + "anthropic_issuer_ttl_seconds": "ttl_seconds", + "anthropic_issuer_signing_key_ref": "signing_key_ref", + } +) +_KEYCLOAK_FIELD_MAP: Final[Mapping[str, str]] = MappingProxyType( + { + "anthropic_keycloak_token_url": "token_url", + "anthropic_keycloak_client_id": "client_id", + "anthropic_keycloak_auth_method": "auth_method", + "anthropic_keycloak_client_secret_ref": "client_secret_ref", + "anthropic_keycloak_scope": "scope", + } +) +_DENIAL_HINT: Final = ( + "Anthropic answers every denied exchange with the same 401; the reason (for example" + " workspace_id_required or jti_reused) is only shown in the Claude Console under" + " Settings > Workload identity, in the rule's authentication history. jti_reused means this" + " identity token was already exchanged once: Anthropic accepts each assertion a single time, so a" + " token file or env var has to rotate before the minted token expires (the rule's" + " token_lifetime_seconds), or switch to the internal issuer or Keycloak source, which mint a" + " fresh assertion per exchange" +) +_WORKSPACE_HINT: Final = ( + "If the federation rule is enabled in more than one workspace, set anthropic_federation_workspace_id" + " (or ANTHROPIC_FEDERATION_WORKSPACE_ID) to the wrkspc_ id of the workspace to mint tokens for, or to 'default'." + " Federation does not read ANTHROPIC_WORKSPACE_ID, which the Bedrock Claude platform provider already uses" +) +_SERVICE_ACCOUNT_HINT: Final = ( + "Anthropic's reference lists service_account_id as required: set anthropic_service_account_id" + " (or ANTHROPIC_SERVICE_ACCOUNT_ID) to the svac_ id the federation rule targets" +) +_MISSING_IDS_HINT: Final = ( + "Copy them from the federation rule's detail page under Settings > Workload identity in the" + " Claude Console, or set ANTHROPIC_FEDERATION_RULE_ID and ANTHROPIC_ORGANIZATION_ID" +) +_ALLOWLIST_HINT: Final = ( + "Identity token files must sit under an allowed credential directory" + " (/var/run/secrets or /run/secrets by default);" + " set LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS to extend the allowlist" +) +_EMPTY_PARAMS: Final[Mapping[str, object]] = MappingProxyType({}) + +_IdentitySourceVariant = TypeVar("_IdentitySourceVariant", bound="InternalIssuerSource | KeycloakSource") + + +class AnthropicWifParams(BaseModel): + model_config = ConfigDict(frozen=True) + + federation_rule_id: str + organization_id: str + service_account_id: str | None = None + workspace_id: str | None = None + assertion_ref: str + assertion_source: Callable[[], str | None] | None = None + + +def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> AnthropicWifParams | None: + if litellm_params is not None and litellm_params.get(_DISABLE_WIF_PARAM) is True: + return None + federation_rule_id: Final = _config_value( + litellm_params, "anthropic_federation_rule_id", "ANTHROPIC_FEDERATION_RULE_ID" + ) + organization_id: Final = _config_value(litellm_params, "anthropic_organization_id", "ANTHROPIC_ORGANIZATION_ID") + if federation_rule_id is None or organization_id is None: + _raise_if_identity_source_configured(litellm_params, federation_rule_id, organization_id) + return None + identity_source: Final = _resolve_identity_source(litellm_params) + if identity_source is None: + return None + assertion_ref, assertion_source = identity_source + return AnthropicWifParams( + federation_rule_id=federation_rule_id, + organization_id=organization_id, + service_account_id=_config_value( + litellm_params, "anthropic_service_account_id", "ANTHROPIC_SERVICE_ACCOUNT_ID" + ), + workspace_id=_config_value( + litellm_params, "anthropic_federation_workspace_id", "ANTHROPIC_FEDERATION_WORKSPACE_ID" + ), + assertion_ref=assertion_ref, + assertion_source=assertion_source, + ) + + +def _resolve_identity_source( + litellm_params: Mapping[str, object] | None, +) -> tuple[str, Callable[[], str] | None] | None: + """Dispatches on ``anthropic_identity_source``. Absent (the default) keeps today's + token_file/env resolution byte-identical, with no ``assertion_source`` closure -- the engine + falls back to its own reader exactly as it does today. A recognized kind builds the matching + frozen config, hashes it into the ``oidc//`` cache-key ref (``identity_source_ref``), + and closes the source's fetch/mint function over it. An unset-but-invalid config (unknown + kind, a missing required field, or a field from the other variant) fails closed here rather + than silently falling back to token_file. A deployment whose params carry a legacy token or + token_file ref stays on legacy resolution even when ``ANTHROPIC_IDENTITY_SOURCE`` names a + fleet-wide kind: the env kind only governs deployments that set no identity params of their own.""" + source_kind: Final = _resolve_source_kind(litellm_params) + if source_kind is None: + legacy_ref: Final = _resolve_assertion_ref(litellm_params) + return (legacy_ref, None) if legacy_ref is not None else None + params: Final[Mapping[str, object]] = MappingProxyType( + {key: value for key, value in (litellm_params or _EMPTY_PARAMS).items() if _is_set(value)} + ) + match source_kind: + case AnthropicIdentitySourceKind.internal_issuer.value: + _reject_foreign_variant_fields(params, foreign_field_map=_KEYCLOAK_FIELD_MAP, chosen_kind=source_kind) + issuer_config: Final = _build_variant(InternalIssuerSource, params, _INTERNAL_ISSUER_FIELD_MAP) + return identity_source_ref(issuer_config), internal_issuer_assertion_source(issuer_config) + case AnthropicIdentitySourceKind.keycloak.value: + _reject_foreign_variant_fields( + params, foreign_field_map=_INTERNAL_ISSUER_FIELD_MAP, chosen_kind=source_kind + ) + keycloak_config: Final = _build_variant(KeycloakSource, params, _KEYCLOAK_FIELD_MAP) + return identity_source_ref(keycloak_config), keycloak_assertion_source(keycloak_config) + case _: + _raise_unknown_source_kind(source_kind) + + +def _raise_unknown_source_kind(source_kind: str) -> NoReturn: + raise litellm.AuthenticationError( + message=( + f"{_IDENTITY_SOURCE_PARAM} must be one of " + f"{', '.join(kind.value for kind in AnthropicIdentitySourceKind)}; got {source_kind!r}" + ), + llm_provider="anthropic", + model="", + ) + + +def _raise_if_identity_source_configured( + litellm_params: Mapping[str, object] | None, federation_rule_id: str | None, organization_id: str | None +) -> None: + """A configured identity source is an explicit request to federate, so a missing rule or + organization id fails closed with the ids named, rather than silently skipping federation + and surfacing later as a missing API key.""" + source_kind: Final = _resolve_source_kind(litellm_params) + if source_kind is None: + return + if source_kind not in {kind.value for kind in AnthropicIdentitySourceKind}: + _raise_unknown_source_kind(source_kind) + missing: Final = tuple( + param + for param, value in ( + ("anthropic_federation_rule_id", federation_rule_id), + ("anthropic_organization_id", organization_id), + ) + if value is None + ) + raise litellm.AuthenticationError( + message=( + f"{_IDENTITY_SOURCE_PARAM} is {source_kind!r}, but {' and '.join(missing)} " + f"{'is' if len(missing) == 1 else 'are'} not set. {_MISSING_IDS_HINT}" + ), + llm_provider="anthropic", + model="", + ) + + +def _resolve_source_kind(litellm_params: Mapping[str, object] | None) -> str | None: + param_kind: Final = _param_str(litellm_params, _IDENTITY_SOURCE_PARAM) + if param_kind is not None: + return param_kind + has_param_legacy_ref: Final = any( + _param_str(litellm_params, key) is not None for key in (_IDENTITY_TOKEN_FILE_PARAM, _IDENTITY_TOKEN_PARAM) + ) + return None if has_param_legacy_ref else _env_str(_IDENTITY_SOURCE_ENV) + + +def _reject_foreign_variant_fields( + litellm_params: Mapping[str, object], foreign_field_map: Mapping[str, str], chosen_kind: str +) -> None: + foreign_keys_present: Final = tuple(param for param in foreign_field_map if param in litellm_params) + if foreign_keys_present: + raise litellm.AuthenticationError( + message=( + f"{_IDENTITY_SOURCE_PARAM} is {chosen_kind!r}, but {', '.join(sorted(foreign_keys_present))} " + "belongs to a different identity source and cannot be set alongside it" + ), + llm_provider="anthropic", + model="", + ) + + +def _build_variant( + model: type[_IdentitySourceVariant], + litellm_params: Mapping[str, object], + field_map: Mapping[str, str], +) -> _IdentitySourceVariant: + fields: Final = MappingProxyType( + {field_map[key]: value for key, value in litellm_params.items() if key in field_map and _is_set(value)} + ) + try: + return model.model_validate(fields) + except ValidationError as e: + # hide_input_in_errors=True on both variant models keeps a secret pasted into the + # wrong field (e.g. a client_secret typed as signing_key_ref) out of str(e). + raise litellm.AuthenticationError( + message=f"Invalid {_IDENTITY_SOURCE_PARAM} configuration: {e}", + llm_provider="anthropic", + model="", + ) from e + + +@dataclass(frozen=True, slots=True) +class ExportedJwks: + document: str + + +@dataclass(frozen=True, slots=True) +class NotAnInternalIssuerCredential: + required_param: str + required_value: str + + +@dataclass(frozen=True, slots=True) +class UnbuildableIdentitySource: + message: str + + +AnthropicJwksExport = ExportedJwks | NotAnInternalIssuerCredential | UnbuildableIdentitySource + + +def anthropic_internal_issuer_jwks(credential_values: Mapping[str, object]) -> AnthropicJwksExport: + """Derive the public JWKS a stored anthropic credential publishes to its federation issuer. + The private signing key stays in this process; only the derived public document comes back.""" + if credential_values.get(_IDENTITY_SOURCE_PARAM) != AnthropicIdentitySourceKind.internal_issuer.value: + return NotAnInternalIssuerCredential( + required_param=_IDENTITY_SOURCE_PARAM, + required_value=AnthropicIdentitySourceKind.internal_issuer.value, + ) + try: + issuer_source: Final = _build_variant(InternalIssuerSource, credential_values, _INTERNAL_ISSUER_FIELD_MAP) + return ExportedJwks(internal_issuer_jwks_document(issuer_source)) + except (litellm.AuthenticationError, ValueError) as e: + return UnbuildableIdentitySource(str(e)) + + +def build_anthropic_wif_spec(params: AnthropicWifParams, api_base: str) -> TokenExchangeSpec: + return TokenExchangeSpec( + token_url=api_base.rstrip("/") + ANTHROPIC_TOKEN_EXCHANGE_PATH, + assertion_ref=params.assertion_ref, + assertion_field="assertion", + static_body=MappingProxyType( + { + name: value + for name, value in ( + ("grant_type", _JWT_BEARER_GRANT_TYPE), + ("federation_rule_id", params.federation_rule_id), + ("organization_id", params.organization_id), + ("service_account_id", params.service_account_id), + ("workspace_id", params.workspace_id), + ) + if value is not None + } + ), + body_encoding="json", + request_headers=MappingProxyType({}), + assertion_source=params.assertion_source, + cache_key_identity=( + params.federation_rule_id, + params.organization_id, + params.service_account_id or "", + params.workspace_id or "", + ), + ) + + +def get_anthropic_wif_token( + litellm_params: Mapping[str, object] | None, + api_base: str | None, + model: str, + engine: JwtBearerTokenExchangeEngine = default_token_exchange_engine, +) -> str | None: + params: Final = resolve_anthropic_wif_params(litellm_params) + if params is None: + return None + exchange_base: Final = resolve_anthropic_base(api_base) + _raise_if_exchange_host_untrusted(exchange_base, model) + result: Final = engine.get_token(build_anthropic_wif_spec(params, exchange_base)) + return _token_from_result(result, model, params) + + +async def aget_anthropic_wif_token( + litellm_params: Mapping[str, object] | None, + api_base: str | None, + model: str, + engine: JwtBearerTokenExchangeEngine = default_token_exchange_engine, +) -> str | None: + params: Final = resolve_anthropic_wif_params(litellm_params) + if params is None: + return None + exchange_base: Final = resolve_anthropic_base(api_base) + _raise_if_exchange_host_untrusted(exchange_base, model) + result: Final = await engine.aget_token(build_anthropic_wif_spec(params, exchange_base)) + return _token_from_result(result, model, params) + + +def _token_from_result(result: ExchangeResult, model: str, params: AnthropicWifParams) -> str: + match result: + case MintedToken(): + return result.access_token.get_secret_value() + case _: + _raise_anthropic_wif_error( + result, + model=model, + workspace_id_set=params.workspace_id is not None, + service_account_id_set=params.service_account_id is not None, + ) + + +def resolve_anthropic_base(api_base: str | None) -> str: + """The base every Anthropic tier derives its URLs from: the deployment api_base when set, + else ``ANTHROPIC_API_BASE`` / ``ANTHROPIC_BASE_URL``, else Anthropic's host, with trailing + slashes and chat-appended ``/v1/messages`` suffixes stripped, so the token URL, the cache key + and the count-tokens URL all agree for the same deployment.""" + return anthropic_base_without_chat_suffix(api_base or _resolve_default_api_base()) + + +def _allowlisted_authority(entry: str) -> tuple[str, int | None] | None: + """One allowlist entry as ``(host, port)``. The port stays ``None`` unless the entry spells one + out, so ``gateway.internal`` trusts that host on every port while ``gateway.internal:8443`` + trusts only 8443.""" + parts: Final = urlsplit(entry if "://" in entry else f"//{entry}") + try: + port: Final = parts.port + except ValueError: + return None + return (parts.hostname, port) if parts.hostname else None + + +def _exchange_authority(exchange_base: str) -> tuple[str, int | None]: + """The host and port an exchange would actually reach, filling in the scheme's default port so + an operator who wrote ``api.anthropic.com:443`` still matches ``https://api.anthropic.com``.""" + parts: Final = urlsplit(exchange_base) + try: + port: Final = parts.port + except ValueError: + return "", None + return (parts.hostname or "").lower(), port if port is not None else _SCHEME_DEFAULT_PORTS.get(parts.scheme) + + +def _trusted_exchange_authorities() -> frozenset[tuple[str, int | None]]: + """Authorities a federated exchange may reach: Anthropic's own, plus whatever the operator put in + the environment. Comma separated, case folded, each entry a URL, a bare host, or ``host:port``.""" + configured: Final = os.getenv(_TRUSTED_EXCHANGE_HOSTS_ENV) or "" + entries: Final = (_allowlisted_authority(entry.strip()) for entry in configured.split(",") if entry.strip()) + return frozenset(chain(((_DEFAULT_TRUSTED_EXCHANGE_HOST, None),), (entry for entry in entries if entry))) + + +def _raise_if_exchange_host_untrusted(exchange_base: str, model: str) -> None: + """The federated exchange refuses any authority the operator has not vouched for, whatever wrote + the deployment's api_base. Exact host match, never a substring: ``api.anthropic.com.evil.test`` + contains the real host and must not pass. An entry naming a port trusts that port alone, so a + second process on another port of an allowed host is refused.""" + host, port = _exchange_authority(exchange_base) + if host and any( + host == allowed_host and allowed_port in (None, port) + for allowed_host, allowed_port in _trusted_exchange_authorities() + ): + return + refused: Final = f"{host}:{port}" if host and port is not None else host + raise litellm.AuthenticationError( + message=( + f"Anthropic workload identity federation refused to use host {refused or exchange_base!r}. " + f"A federated exchange sends the workload's identity token to this host and presents the " + f"minted token to it, so only {_DEFAULT_TRUSTED_EXCHANGE_HOST} is trusted by default. To " + f"use a private Anthropic-compatible gateway, add its host, or host:port to pin the port, " + f"to the {_TRUSTED_EXCHANGE_HOSTS_ENV} environment variable (comma separated); that is a " + f"decision to trust it with org-scoped credentials, so it is deliberately server-owned " + f"and cannot be set through the model or credential APIs" + ), + llm_provider="anthropic", + model=model, + ) + + +@lru_cache(maxsize=_SHADOWED_DEPLOYMENT_WARNING_CAP) +def _warn_static_credential_shadows_federation(model: str, configured_rule_id: str | None) -> None: + """Memoized so a shadowed deployment says this once rather than once per request. + + The environment fallback is resolved in here rather than by the caller so it too costs one + secret-manager read per deployment: every static-key Anthropic call reaches this, and a + per-request read of a rule id almost nobody sets is an ERROR log with a traceback per call on + the deployments that shadow nothing. + """ + if configured_rule_id is None and _env_str("ANTHROPIC_FEDERATION_RULE_ID") is None: + return + verbose_logger.warning( + "Anthropic deployment %s is configured for workload identity federation, but a static " + "ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN is set and takes precedence, so every call bills " + "that credential and no federated token is minted. Unset it to federate.", + model or "(unnamed)", + ) + + +def warn_if_static_credential_shadows_federation(litellm_params: Mapping[str, object] | None, model: str) -> None: + """A process-wide static credential outranks federation everywhere in the provider, which is the + Anthropic SDK's own precedence. An operator who configured federation and left a key behind would + otherwise get no signal at all that none of their calls are federated.""" + if litellm_params is not None and litellm_params.get(_DISABLE_WIF_PARAM) is True: + return + _warn_static_credential_shadows_federation(model, _param_str(litellm_params, "anthropic_federation_rule_id")) + + +def _resolve_default_api_base() -> str: + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + return AnthropicModelInfo.get_api_base(None) or _DEFAULT_API_BASE + + +def anthropic_base_without_chat_suffix(base: str) -> str: + """A deployment base with its chat-surface suffix removed, so the token URL and model + discovery both derive from the same value whatever form the operator configured.""" + parts: Final = urlsplit(base) + if not parts.scheme or not parts.netloc: + return base.rstrip("/") + return urlunsplit((parts.scheme, parts.netloc, _strip_path_suffixes(parts.path), "", "")) + + +def _strip_path_suffixes(path: str) -> str: + """Drop the chat-surface suffixes a deployment base may carry, so every tier derives the same + token URL. Each pass removes at most one suffix, so the loop is bounded by the segment count.""" + trimmed = path.rstrip("/") # rebind-ok: fixed-point strip, one suffix per pass + while True: + shortened = next( + (trimmed.removesuffix(suffix) for suffix in _CHAT_BASE_SUFFIXES if trimmed.endswith(suffix)), + trimmed, + ) + if shortened == trimmed: + return trimmed + # Re-strip: a doubled suffix leaves a trailing slash that would stop the next match. + trimmed = shortened.rstrip("/") + + +def _config_value(litellm_params: Mapping[str, object] | None, param_key: str, env_name: str) -> str | None: + return _param_str(litellm_params, param_key) or _env_str(env_name) + + +def _is_set(value: object) -> bool: + return value is not None and value != "" + + +def _param_str(litellm_params: Mapping[str, object] | None, key: str) -> str | None: + if litellm_params is None: + return None + value: Final = litellm_params.get(key) + return value if isinstance(value, str) and value else None + + +def _env_str(name: str) -> str | None: + from litellm.secret_managers.main import get_secret_str + + value: Final = get_secret_str(name) + return value if isinstance(value, str) and value else None + + +def _resolve_assertion_ref(litellm_params: Mapping[str, object] | None) -> str | None: + file_param: Final = _param_str(litellm_params, _IDENTITY_TOKEN_FILE_PARAM) + if file_param is not None: + return f"oidc/file/{file_param}" + inline_param: Final = _param_str(litellm_params, _IDENTITY_TOKEN_PARAM) + if inline_param is not None: + return _validated_inline_ref(inline_param) + file_env: Final = _env_str("ANTHROPIC_IDENTITY_TOKEN_FILE") + if file_env is not None: + return f"oidc/file/{file_env}" + if _env_str(_INLINE_ENV_VAR) is not None: + return f"oidc/env/{_INLINE_ENV_VAR}" + return None + + +def _validated_inline_ref(value: str) -> str: + if value.startswith(_ACCEPTED_REF_PREFIX) and not value.startswith(_REJECTED_REF_PREFIX): + return value + raise litellm.AuthenticationError( + message=( + "anthropic_identity_token must be an oidc/ secret reference such as oidc/env/VAR_NAME," + " oidc/file//absolute/path, oidc/github/, or oidc/google/." + " Raw identity tokens and oidc/env_path/ references are not accepted;" + " to pass a token directly, export it and reference it as oidc/env/VAR_NAME" + ), + llm_provider="anthropic", + model="", + ) + + +def _raise_anthropic_wif_error( + error: ExchangeError, model: str, workspace_id_set: bool, service_account_id_set: bool +) -> NoReturn: + detail: Final = _error_detail( + error, workspace_id_set=workspace_id_set, service_account_id_set=service_account_id_set + ) + raise litellm.AuthenticationError( + message=f"Anthropic workload identity federation failed. {detail}", + llm_provider="anthropic", + model=model, + ) + + +def _denial_hints(workspace_id_set: bool, service_account_id_set: bool) -> str: + hints: Final = ( + _DENIAL_HINT, + "" if workspace_id_set else _WORKSPACE_HINT, + "" if service_account_id_set else _SERVICE_ACCOUNT_HINT, + ) + return " " + ". ".join(hint for hint in hints if hint) + + +def _error_detail(error: ExchangeError, workspace_id_set: bool, service_account_id_set: bool) -> str: + match error: + case AssertionSourceError() if error.kind == "disallowed_path": + return f"Could not read the OIDC identity token from {error.source_ref}. {_ALLOWLIST_HINT}" + case AssertionSourceError(): + base: Final = f"Could not obtain the OIDC identity token ({error.kind}) from {error.source_ref}" + return f"{base}. {error.detail}" if error.detail else base + case InsecureTokenUrl(): + return f"The token endpoint must use https; refusing to send the identity token to host {error.host!r}" + case TokenEndpointError() if error.status_code == 401: + hints: Final = _denial_hints(workspace_id_set, service_account_id_set) + return f"The token endpoint returned HTTP 401: {error.redacted_body}{hints}" + case TokenEndpointError(): + return f"The token endpoint returned HTTP {error.status_code}: {error.redacted_body}" + case TokenTransportError(): + return f"Could not reach the token endpoint: {error.detail}" + case MalformedTokenResponse(): + return f"The token endpoint returned an unusable response: {error.detail}" + case _: + assert_never(error) diff --git a/litellm/llms/azure_ai/embed/handler.py b/litellm/llms/azure_ai/embed/handler.py index c65edbf56e6..8037be8abeb 100644 --- a/litellm/llms/azure_ai/embed/handler.py +++ b/litellm/llms/azure_ai/embed/handler.py @@ -1,3 +1,4 @@ +from collections.abc import Mapping from typing import Final from urllib.parse import urlsplit, urlunsplit @@ -218,6 +219,7 @@ class AzureAIEmbedding(OpenAIChatCompletion): aembedding=None, max_retries: int | None = None, shared_session=None, + litellm_params: Mapping[str, object] | None = None, ) -> EmbeddingResponse: """ - Separate image url from text diff --git a/litellm/llms/base_llm/anthropic_messages/transformation.py b/litellm/llms/base_llm/anthropic_messages/transformation.py index 101a5e6c58c..5758d156a61 100644 --- a/litellm/llms/base_llm/anthropic_messages/transformation.py +++ b/litellm/llms/base_llm/anthropic_messages/transformation.py @@ -41,6 +41,29 @@ class BaseAnthropicMessagesConfig(ABC): """ return headers, api_base + async def avalidate_anthropic_messages_environment( + self, + headers: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + model: str, + messages: list[Any], # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + optional_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + litellm_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + api_key: str | None = None, + api_base: str | None = None, + ) -> tuple[dict, str | None]: # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract + """Async counterpart used by the async handler. The default delegates to the + sync implementation; providers whose sync path can block the event loop + (e.g. a WIF token exchange) override this.""" + return self.validate_anthropic_messages_environment( + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=api_key, + api_base=api_base, + ) + @abstractmethod def get_complete_url( self, diff --git a/litellm/llms/base_llm/auth/__init__.py b/litellm/llms/base_llm/auth/__init__.py new file mode 100644 index 00000000000..291e1492c98 --- /dev/null +++ b/litellm/llms/base_llm/auth/__init__.py @@ -0,0 +1,99 @@ +from litellm.llms.base_llm.auth.client_credentials import ( + SecretReader, + fetch_keycloak_assertion, + keycloak_assertion_source, +) +from litellm.llms.base_llm.auth.identity_source import ( + AnthropicIdentitySourceConfig, + AnthropicIdentitySourceKind, + InternalIssuerSource, + KeycloakSource, + identity_source_config_adapter, + identity_source_ref, +) +from litellm.llms.base_llm.auth.internal_issuer import ( + SigningKeyReader, + internal_issuer_assertion_source, + internal_issuer_jwks_document, + mint_internal_issuer_assertion, +) +from litellm.llms.base_llm.auth.jwt_signing import ( + ALG, + build_jwk, + build_jwks, + jwks_document_json, + load_es256_private_key, + rfc7638_thumbprint, + sign_es256_jwt, +) +from litellm.llms.base_llm.auth.token_exchange import ( + ADVISORY_REFRESH_BACKOFF_SECONDS, + ADVISORY_REFRESH_SECONDS, + MANDATORY_REFRESH_SECONDS, + MAX_ASSERTION_BYTES, + MAX_RESPONSE_BYTES, + JwtBearerTokenExchangeEngine, + default_token_exchange_engine, + redact_oauth_error_body, + validate_token_endpoint_url, +) +from litellm.llms.base_llm.auth.types import ( + AssertionReader, + AssertionSource, + AssertionSourceError, + BodyEncoding, + ExchangeError, + ExchangeResult, + InsecureTokenUrl, + MalformedTokenResponse, + MintedToken, + SyncTokenPoster, + TokenEndpointError, + TokenExchangeSpec, + TokenTransportError, +) + +__all__ = ( + "ADVISORY_REFRESH_BACKOFF_SECONDS", + "ADVISORY_REFRESH_SECONDS", + "ALG", + "MANDATORY_REFRESH_SECONDS", + "MAX_ASSERTION_BYTES", + "MAX_RESPONSE_BYTES", + "AnthropicIdentitySourceConfig", + "AnthropicIdentitySourceKind", + "AssertionReader", + "AssertionSource", + "AssertionSourceError", + "BodyEncoding", + "ExchangeError", + "ExchangeResult", + "InsecureTokenUrl", + "InternalIssuerSource", + "JwtBearerTokenExchangeEngine", + "KeycloakSource", + "MalformedTokenResponse", + "MintedToken", + "SecretReader", + "SigningKeyReader", + "SyncTokenPoster", + "TokenEndpointError", + "TokenExchangeSpec", + "TokenTransportError", + "build_jwk", + "build_jwks", + "default_token_exchange_engine", + "fetch_keycloak_assertion", + "identity_source_config_adapter", + "identity_source_ref", + "internal_issuer_assertion_source", + "internal_issuer_jwks_document", + "jwks_document_json", + "keycloak_assertion_source", + "load_es256_private_key", + "mint_internal_issuer_assertion", + "redact_oauth_error_body", + "rfc7638_thumbprint", + "sign_es256_jwt", + "validate_token_endpoint_url", +) diff --git a/litellm/llms/base_llm/auth/client_credentials.py b/litellm/llms/base_llm/auth/client_credentials.py new file mode 100644 index 00000000000..bf86dc6508f --- /dev/null +++ b/litellm/llms/base_llm/auth/client_credentials.py @@ -0,0 +1,225 @@ +"""Fetches a fresh RFC 6749 client_credentials assertion for Anthropic's ``keycloak`` identity +source: LiteLLM authenticates to Keycloak as its own confidential client and presents the +resulting ``access_token`` as the workload assertion (Phase 1 decision 2). + +The client secret is the operator-supplied pointer at ``KeycloakSource.client_secret_ref``, +resolved the same way every other WIF secret pointer already is (env, a Credential, or whatever +secret manager ``litellm.secret_manager_client`` is globally configured to, Vault included). +Every fetch is a fresh HTTP POST; nothing here caches a fetched token, since the outer +token-exchange engine already caches the Anthropic token it buys with one -- see decision 2's +"no Keycloak-side cache" ruling. +""" + +import base64 +import threading +from collections.abc import Callable, Mapping +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, TypeAlias +from urllib.parse import quote, quote_plus, urlencode + +import httpx +from pydantic import BaseModel, SecretStr, ValidationError +from typing_extensions import assert_never + +from litellm.llms.base_llm.auth.identity_source import KeycloakSource, ref_for_error_message +from litellm.llms.base_llm.auth.token_exchange import ( + MAX_RESPONSE_BYTES, + endpoint_url_for_error_message, + redact_oauth_error_body, + require_posted_response, + validate_token_endpoint_url, +) +from litellm.llms.base_llm.auth.types import InsecureTokenUrl, SyncTokenPoster + +if TYPE_CHECKING: + from litellm.llms.custom_httpx.http_handler import HTTPHandler + +SecretReader: TypeAlias = Callable[[str], str | None] + +_GRANT_TYPE: Final = "client_credentials" +_TIMEOUT_SECONDS: Final = 30.0 +_FORM_CONTENT_TYPE: Final = "application/x-www-form-urlencoded" + + +class _ClientCredentialsResponse(BaseModel): + access_token: str + + +def _default_secret_reader(ref: str) -> str | None: + from litellm.secret_managers.main import get_secret_str + + return get_secret_str(ref) + + +def _new_keycloak_handler() -> "HTTPHandler": + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + return HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0), follow_redirects=False) + + +class _HttpxSyncKeycloakPoster: + """Dedicated HTTPHandler for the Keycloak token POST: no ``logging_obj`` (so litellm's + request/response logging never sees the client secret or the fetched token), redirects + disabled. A separate instance from the outer engine's own poster, since this is a genuinely + new HTTP call site whose no-logging guarantee must be built here, not assumed inherited.""" + + def __init__(self, handler_factory: Callable[[], "HTTPHandler"] = _new_keycloak_handler) -> None: + self._lock: Final = threading.Lock() + self._handler_factory: Final = handler_factory + self._handler: HTTPHandler | None = None + + def _handler_instance(self) -> "HTTPHandler": + with self._lock: + if self._handler is None: + self._handler = self._handler_factory() + return self._handler + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + try: + response: Final[httpx.Response | None] = self._handler_instance().post( # pyright: ignore[reportUnknownMemberType] # HTTPHandler.post is legacy-untyped; the result is validated below + url, + content=content, + headers=dict(headers), + timeout=timeout, + ) + except httpx.HTTPStatusError as e: + return e.response + return require_posted_response(response, "keycloak token endpoint") + + +_DEFAULT_POSTER: Final[SyncTokenPoster] = _HttpxSyncKeycloakPoster() + + +def _form_encode(value: str) -> str: + """RFC 6749 Appendix B before RFC 6749 2.3.1's base64: application/x-www-form-urlencoded + with spaces as ``%20`` rather than ``+``, else a reserved character (":", "+", "%", " ") in + the id or secret corrupts the credential the far side decodes back out of Basic auth.""" + return quote(value, safe="") + + +def _basic_auth_header(client_id: str, client_secret: str) -> str: + encoded_pair: Final = f"{_form_encode(client_id)}:{_form_encode(client_secret)}" + return "Basic " + base64.b64encode(encoded_pair.encode()).decode("ascii") + + +def _prepared_request(config: KeycloakSource, client_secret: str) -> tuple[bytes, Mapping[str, str]]: + scope_field: Final[Mapping[str, str]] = ( + MappingProxyType({"scope": config.scope}) if config.scope else MappingProxyType({}) + ) + match config.auth_method: + case "client_secret_basic": + return ( + urlencode(MappingProxyType({"grant_type": _GRANT_TYPE, **scope_field})).encode(), + MappingProxyType( + { + "content-type": _FORM_CONTENT_TYPE, + "authorization": _basic_auth_header(config.client_id, client_secret), + } + ), + ) + case "client_secret_post": + return ( + urlencode( + MappingProxyType( + { + "grant_type": _GRANT_TYPE, + "client_id": config.client_id, + "client_secret": client_secret, + **scope_field, + } + ) + ).encode(), + MappingProxyType({"content-type": _FORM_CONTENT_TYPE}), + ) + case _: + assert_never(config.auth_method) + + +def _resolve_client_secret(config: KeycloakSource, secret_reader: SecretReader) -> str: + secret: Final = secret_reader(config.client_secret_ref) + if not secret: + raise ValueError(f"keycloak client secret {ref_for_error_message(config.client_secret_ref)} could not be read") + return secret + + +def _wire_forms_of_secret(config: KeycloakSource, client_secret: str) -> tuple[SecretStr, ...]: + """Every shape the secret leaves this process in, so an echo of any of them is caught. + + Neither grant sends the secret verbatim. client_secret_basic base64s ``id:secret``, which + decodes straight back to it, and client_secret_post percent-escapes it. An endpoint echoing + either shape hands over reversible material a raw comparison would miss. + """ + raw: Final = SecretStr(client_secret) + match config.auth_method: + case "client_secret_basic": + encoded_pair: Final = f"{_form_encode(config.client_id)}:{_form_encode(client_secret)}" + return (raw, SecretStr(base64.b64encode(encoded_pair.encode()).decode("ascii"))) + case "client_secret_post": + # urlencode escapes reserved characters and writes a space as "+", so a secret + # containing either leaves in a shape the raw comparison would not recognise coming + # back. quote_plus is what urlencode itself applies. + return (raw, SecretStr(quote_plus(client_secret))) + case _: + assert_never(config.auth_method) + + +def _endpoint_error_message(config: KeycloakSource, response: httpx.Response, client_secret: str) -> str: + endpoint_error: Final = redact_oauth_error_body( + response.status_code, response.text, _wire_forms_of_secret(config, client_secret) + ) + return ( + f"keycloak token endpoint {endpoint_url_for_error_message(config.token_url)} " + f"returned HTTP {endpoint_error.status_code}: {endpoint_error.redacted_body}" + ) + + +def _parse_success_body(response: httpx.Response) -> str: + if len(response.content) > MAX_RESPONSE_BYTES: + raise ValueError("keycloak token response exceeded the size cap") + try: + parsed: Final = _ClientCredentialsResponse.model_validate_json(response.content) + except ValidationError as e: + raise ValueError("keycloak token response failed schema validation") from e + token: Final = parsed.access_token.strip() + if not token: + raise ValueError("keycloak token response carried an empty access_token") + return token + + +def fetch_keycloak_assertion( + config: KeycloakSource, + *, + poster: SyncTokenPoster = _DEFAULT_POSTER, + secret_reader: SecretReader = _default_secret_reader, +) -> str: + """POSTs one fresh client_credentials grant and returns the resulting ``access_token`` as the + workload assertion; the caller must not cache the result -- see the module docstring.""" + match validate_token_endpoint_url(config.token_url): + case InsecureTokenUrl(host=host): + raise ValueError(f"keycloak token_url must use https; refusing to send the client secret to host {host!r}") + case _: + pass + client_secret: Final = _resolve_client_secret(config, secret_reader) + content, headers = _prepared_request(config, client_secret) + try: + response: Final = poster.post(config.token_url, content=content, headers=headers, timeout=_TIMEOUT_SECONDS) + except Exception as e: # noqa: BLE001 # injected posters may raise beyond httpx; every failure becomes a ValueError + raise ValueError( + f"could not reach the keycloak token endpoint {endpoint_url_for_error_message(config.token_url)}: " + f"{type(e).__name__}" + ) from e + if not 200 <= response.status_code < 300: + raise ValueError(_endpoint_error_message(config, response, client_secret)) + return _parse_success_body(response) + + +def keycloak_assertion_source( + config: KeycloakSource, + *, + poster: SyncTokenPoster = _DEFAULT_POSTER, + secret_reader: SecretReader = _default_secret_reader, +) -> Callable[[], str]: + """A zero-arg closure that fetches fresh on every call: the shape an ``oidc/keycloak/...`` + ref dispatches to once wired into ``TokenExchangeSpec.assertion_source`` (Phase 1 decision 7) + -- the caller parses the config and closes this function over it, with no registry involved.""" + return lambda: fetch_keycloak_assertion(config, poster=poster, secret_reader=secret_reader) diff --git a/litellm/llms/base_llm/auth/identity_source.py b/litellm/llms/base_llm/auth/identity_source.py new file mode 100644 index 00000000000..f8f02249b3c --- /dev/null +++ b/litellm/llms/base_llm/auth/identity_source.py @@ -0,0 +1,76 @@ +"""Tagged-union identity-source configs for Anthropic workload identity federation, beyond the +existing token_file/env resolver in ``litellm/llms/anthropic/wif.py``. + +Each variant only ever carries secret *pointer names* (``signing_key_ref``, ``client_secret_ref``), +never a resolved secret value, so ``identity_source_ref`` can safely hash a variant into the short, +content-derived ``oidc//`` string used elsewhere as a get_secret ref, a token-exchange +cache-key discriminator, and an operator-facing error pointer. +""" + +import hashlib +from enum import Enum +from typing import Annotated, Final, Literal, TypeAlias + +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter + +_REF_HASH_HEX_LENGTH: Final = 16 +_MAX_TTL_SECONDS: Final = 3600 +_DEFAULT_TTL_SECONDS: Final = 300 + + +class AnthropicIdentitySourceKind(str, Enum): + internal_issuer = "internal_issuer" + keycloak = "keycloak" + + +class InternalIssuerSource(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", hide_input_in_errors=True) + + kind: Literal[AnthropicIdentitySourceKind.internal_issuer] = AnthropicIdentitySourceKind.internal_issuer + issuer_url: str + subject: str + audience: str | None = None + ttl_seconds: Annotated[int, Field(gt=0, le=_MAX_TTL_SECONDS)] = _DEFAULT_TTL_SECONDS + signing_key_ref: str + + +class KeycloakSource(BaseModel): + model_config = ConfigDict(frozen=True, extra="forbid", hide_input_in_errors=True) + + kind: Literal[AnthropicIdentitySourceKind.keycloak] = AnthropicIdentitySourceKind.keycloak + token_url: str + client_id: str + auth_method: Literal["client_secret_basic", "client_secret_post"] = "client_secret_basic" + client_secret_ref: str + scope: str | None = None + + +AnthropicIdentitySourceConfig: TypeAlias = Annotated[InternalIssuerSource | KeycloakSource, Field(discriminator="kind")] +identity_source_config_adapter: Final = TypeAdapter[AnthropicIdentitySourceConfig](AnthropicIdentitySourceConfig) + + +def identity_source_ref(config: AnthropicIdentitySourceConfig) -> str: + """``oidc//``: a short, secret-free pointer, stable for identical config and rolling + whenever any field does, including a ``*_ref`` pointer NAME (never the secret it points to).""" + digest: Final = hashlib.sha256(config.model_dump_json().encode()).hexdigest()[:_REF_HASH_HEX_LENGTH] + return f"oidc/{config.kind.value}/{digest}" + + +_POINTER_REF_PREFIXES: Final = ( + "oidc/", + "os.environ/", + "hashicorp_vault/", + "aws_secret_manager/", + "google_secret_manager/", +) + + +def ref_for_error_message(ref: str) -> str: + """A ``*_ref`` rendered for an operator-facing error. + + Naming the pointer is deliberate: it is what tells an operator which setting failed to + resolve. But these fields only ever fail to resolve when what was written is not a pointer, + and an operator who pasted the secret itself has made the field's value the secret. So the + value is echoed only when it is recognizably a pointer, and withheld otherwise. + """ + return ref if ref.startswith(_POINTER_REF_PREFIXES) else "" diff --git a/litellm/llms/base_llm/auth/internal_issuer.py b/litellm/llms/base_llm/auth/internal_issuer.py new file mode 100644 index 00000000000..f1029de0757 --- /dev/null +++ b/litellm/llms/base_llm/auth/internal_issuer.py @@ -0,0 +1,86 @@ +"""Mints a self-issued workload assertion for Anthropic's ``internal_issuer`` identity source: +LiteLLM signs its own short-lived ES256 JWT instead of reading one from a mounted OIDC file. + +Signing custody is the operator-supplied PEM at ``InternalIssuerSource.signing_key_ref``, +resolved the same way every other WIF secret pointer already is (env, a Credential, or +whatever secret manager ``litellm.secret_manager_client`` is globally configured to, Vault +included) -- see Phase 1 decision 1. Every mint is fresh; nothing here caches a minted JWT, +since the outer token-exchange engine already caches the Anthropic token it buys with one. +""" + +import time +import uuid +from collections.abc import Callable, Mapping +from types import MappingProxyType +from typing import Final, TypeAlias + +from litellm.llms.base_llm.auth.identity_source import InternalIssuerSource, ref_for_error_message +from litellm.llms.base_llm.auth.jwt_signing import jwks_document_json, sign_es256_jwt + +SigningKeyReader: TypeAlias = Callable[[str], str | None] + + +def _default_signing_key_reader(ref: str) -> str | None: + from litellm.secret_managers.main import get_secret_str + + return get_secret_str(ref) + + +def _claims(config: InternalIssuerSource, issued_at: int) -> Mapping[str, object]: + return MappingProxyType( + { + key: value + for key, value in ( + ("sub", config.subject), + ("iss", config.issuer_url), + ("aud", config.audience), + ("iat", issued_at), + ("exp", issued_at + config.ttl_seconds), + ("jti", str(uuid.uuid4())), + ) + if value is not None + } + ) + + +def _resolve_signing_key(config: InternalIssuerSource, key_reader: SigningKeyReader) -> str: + pem: Final = key_reader(config.signing_key_ref) + if not pem: + raise ValueError( + f"internal_issuer signing key {ref_for_error_message(config.signing_key_ref)} could not be read" + ) + return pem + + +def mint_internal_issuer_assertion( + config: InternalIssuerSource, + *, + key_reader: SigningKeyReader = _default_signing_key_reader, + clock: Callable[[], float] = time.time, +) -> str: + """Signs one fresh, short-lived assertion; the caller must not cache the result, since a + cached copy would defeat the point of re-minting on every exchange.""" + pem: Final = _resolve_signing_key(config, key_reader) + return sign_es256_jwt(pem, _claims(config, issued_at=int(clock()))) + + +def internal_issuer_assertion_source( + config: InternalIssuerSource, + *, + key_reader: SigningKeyReader = _default_signing_key_reader, + clock: Callable[[], float] = time.time, +) -> Callable[[], str]: + """A zero-arg closure that mints fresh on every call: the shape an ``oidc/internal_issuer/...`` + ref dispatches to once wired into ``TokenExchangeSpec.assertion_source`` (Phase 1 decision 7) + -- the caller parses the config and closes this function over it, with no registry involved.""" + return lambda: mint_internal_issuer_assertion(config, key_reader=key_reader, clock=clock) + + +def internal_issuer_jwks_document( + config: InternalIssuerSource, + *, + key_reader: SigningKeyReader = _default_signing_key_reader, +) -> str: + """The operator-facing JWKS export, resolved from a configured identity source rather than + a raw PEM in hand -- the JSON document to register as Anthropic's inline federation issuer.""" + return jwks_document_json(_resolve_signing_key(config, key_reader)) diff --git a/litellm/llms/base_llm/auth/jwt_signing.py b/litellm/llms/base_llm/auth/jwt_signing.py new file mode 100644 index 00000000000..b2079e84a35 --- /dev/null +++ b/litellm/llms/base_llm/auth/jwt_signing.py @@ -0,0 +1,115 @@ +"""ES256 JWT signing primitives for Anthropic workload identity federation's +``internal_issuer`` identity source (see ``identity_source.InternalIssuerSource``). + +Pure functions over an already-resolved PEM string: no I/O, no secret-manager awareness, no +caching. Given the signing key at, say, $ISSUER_SIGNING_KEY_PEM, an operator publishes the +JWKS document Anthropic's inline federation issuer needs with one line: + + python -c "from litellm.llms.base_llm.auth.jwt_signing import jwks_document_json; \\ + import os; print(jwks_document_json(os.environ['ISSUER_SIGNING_KEY_PEM']))" +""" + +import base64 +import hashlib +import json +from collections.abc import Mapping +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, TypeAlias + +if TYPE_CHECKING: + from cryptography.hazmat.primitives.asymmetric import ec + +ALG: Final = "ES256" +MISSING_SIGNING_DEPENDENCIES_MESSAGE: Final = ( + "the internal_issuer identity source needs PyJWT and cryptography, which a base litellm install " + "does not include: pip install 'litellm[proxy]'" +) +_JWK_CURVE_NAME: Final = "P-256" +_JWK_KEY_TYPE: Final = "EC" +_COORDINATE_BYTE_LENGTH: Final = 32 # P-256 field element width, RFC 7518 6.2.1.2/6.2.1.3 + +Jwk: TypeAlias = Mapping[str, str] +Jwks: TypeAlias = Mapping[str, tuple[Jwk, ...]] + + +def load_es256_private_key(pem: str) -> "ec.EllipticCurvePrivateKey": + """Parses an unencrypted PEM EC private key. Never echoes the key material in an error.""" + try: + from cryptography.hazmat.primitives.asymmetric import ec + from cryptography.hazmat.primitives.serialization import load_pem_private_key + except ImportError as e: + raise ImportError(MISSING_SIGNING_DEPENDENCIES_MESSAGE) from e + try: + key: Final = load_pem_private_key(pem.encode(), password=None) + except (ValueError, TypeError) as e: + raise ValueError("internal_issuer signing key is not a valid unencrypted PEM private key") from e + if not isinstance(key, ec.EllipticCurvePrivateKey) or not isinstance(key.curve, ec.SECP256R1): + raise ValueError( # noqa: TRY004 # the reader classifies ValueError into a readable config error; TypeError would not + "internal_issuer signing key must be an EC P-256 (secp256r1) private key for ES256" + ) + return key + + +def _b64url_coordinate(value: int) -> str: + return base64.urlsafe_b64encode(value.to_bytes(_COORDINATE_BYTE_LENGTH, "big")).rstrip(b"=").decode("ascii") + + +def _jwk_thumbprint_members(public_key: "ec.EllipticCurvePublicKey") -> Jwk: + """RFC 7638 3.2's exact EC member set (crv, kty, x, y) and nothing else: an extra member + here would change the thumbprint and desync it from the ``kid`` published in the JWKS.""" + numbers: Final = public_key.public_numbers() + return MappingProxyType( + { + "crv": _JWK_CURVE_NAME, + "kty": _JWK_KEY_TYPE, + "x": _b64url_coordinate(numbers.x), + "y": _b64url_coordinate(numbers.y), + } + ) + + +def rfc7638_thumbprint(public_key: "ec.EllipticCurvePublicKey") -> str: + """RFC 7638: SHA-256 over the lexicographically member-ordered, whitespace-free JSON + rendering of the thumbprint members, base64url-encoded without padding.""" + canonical: Final = json.dumps( + dict(sorted(_jwk_thumbprint_members(public_key).items())), + separators=(",", ":"), + ) + return base64.urlsafe_b64encode(hashlib.sha256(canonical.encode()).digest()).rstrip(b"=").decode("ascii") + + +def build_jwk(public_key: "ec.EllipticCurvePublicKey", kid: str) -> Jwk: + return MappingProxyType({**_jwk_thumbprint_members(public_key), "use": "sig", "alg": ALG, "kid": kid}) + + +def build_jwks(public_key: "ec.EllipticCurvePublicKey") -> Jwks: + kid: Final = rfc7638_thumbprint(public_key) + return MappingProxyType({"keys": (build_jwk(public_key, kid),)}) + + +def jwks_document_json(pem: str) -> str: + """The operator-facing export: the JSON document to register as Anthropic's inline JWKS. + + ``build_jwks`` returns ``MappingProxyType``/tuple values per this repo's no-mutation + convention; the ``json`` module only knows plain ``dict``/``list``, so those are converted + at this one serialization boundary rather than giving up immutability throughout the module. + """ + key: Final = load_es256_private_key(pem) + jwks: Final = build_jwks(key.public_key()) + return json.dumps( + {"keys": [dict(jwk) for jwk in jwks["keys"]]}, + indent=2, + ) + + +def sign_es256_jwt(pem: str, claims: Mapping[str, object]) -> str: + """Signs ``claims`` with the PEM key, stamping ``kid`` as its RFC 7638 thumbprint so a + verifier can look the signing key up in the published JWKS by ``kid`` alone.""" + try: + import jwt + except ImportError as e: + raise ImportError(MISSING_SIGNING_DEPENDENCIES_MESSAGE) from e + key: Final = load_es256_private_key(pem) + kid: Final = rfc7638_thumbprint(key.public_key()) + headers: Final = {"kid": kid} + return jwt.encode(dict(claims), key, algorithm=ALG, headers=headers) diff --git a/litellm/llms/base_llm/auth/shared_token_store.py b/litellm/llms/base_llm/auth/shared_token_store.py new file mode 100644 index 00000000000..878df1dc6fb --- /dev/null +++ b/litellm/llms/base_llm/auth/shared_token_store.py @@ -0,0 +1,181 @@ +"""Same-host token store for the JWT-bearer exchange engine. + +Anthropic accepts an assertion carrying a ``jti`` once per issuer, so every uvicorn worker that +reads the same projected token file must share the token the first exchange minted instead of +re-sending the same assertion. The engine keys the store by cache key and only reuses a stored +token minted from the assertion it currently holds; a rotated assertion always buys a fresh token. +""" + +import contextlib +import os +import sys +import tempfile +import threading +from collections.abc import Generator +from dataclasses import dataclass +from pathlib import Path +from typing import Final, Protocol + +from pydantic import BaseModel, SecretStr, ValidationError + +from litellm._logging import verbose_logger + +CACHE_DIR_ENV: Final = "LITELLM_TOKEN_EXCHANGE_CACHE_DIR" + + +@dataclass(frozen=True, slots=True) +class StoredToken: + access_token: SecretStr + expires_at_epoch: float | None + assertion_sha256: str + + +class SharedTokenStore(Protocol): + """Every method is best-effort: a store that cannot read, write, or lock degrades to a per-process + cache and never raises into the mint path.""" + + def load(self, key: str) -> StoredToken | None: ... + + def save(self, key: str, token: StoredToken) -> None: ... + + def delete(self, key: str) -> None: ... + + def lock(self, key: str) -> contextlib.AbstractContextManager[None]: ... + + +class _StoredTokenFile(BaseModel): + access_token: str + expires_at_epoch: float | None + assertion_sha256: str + + +def _directory_is_private(directory: Path) -> bool: + try: + directory.mkdir(mode=0o700, exist_ok=True) + stat: Final = directory.stat() + except OSError as e: + verbose_logger.warning("Token exchange cache directory %s is unusable (%s); caching per process", directory, e) + return False + if stat.st_uid != os.getuid() or stat.st_mode & 0o077: + verbose_logger.warning( + "Token exchange cache directory %s must be owned by uid %d with mode 0700; caching per process", + directory, + os.getuid(), + ) + return False + return True + + +def _unlink(path: Path) -> None: + with contextlib.suppress(OSError): + path.unlink() + + +def _write_token_file(directory: Path, key: str, body: bytes) -> None: + """The token is staged in its own file and renamed over the entry, so a reader never sees a + half-written one. Every failure unlinks the staging file, including the buffered write that only + reaches the disk when the handle closes: nothing else sweeps this directory, and that file holds + a token that still works. The rename leaves nothing behind for the unlink to find.""" + descriptor, name = tempfile.mkstemp(dir=directory, prefix=f"{key}.") + os.close(descriptor) + staged: Final = Path(name) + try: + staged.write_bytes(body) + os.replace(staged, directory / f"{key}.json") + finally: + _unlink(staged) + + +class FileTokenStore: + """One ``.json`` (mode 0600) and one ``.lock`` (flock) per identity under a + directory only the proxy's uid can enter; the directory is checked on first use, not at import.""" + + def __init__(self, directory: Path) -> None: + self._directory: Final = directory + self._ready_lock: Final = threading.Lock() + self._ready: bool | None = None + + @property + def directory(self) -> Path: + return self._directory + + def _usable(self) -> bool: + with self._ready_lock: + if self._ready is None: + self._ready = _directory_is_private(self._directory) + return self._ready + + def load(self, key: str) -> StoredToken | None: + if not self._usable(): + return None + try: + raw: Final = (self._directory / f"{key}.json").read_bytes() + parsed: Final = _StoredTokenFile.model_validate_json(raw) + except FileNotFoundError: + return None + except (OSError, ValidationError) as e: + verbose_logger.debug("Ignoring unreadable token exchange cache entry: %s", e) + return None + return StoredToken( + access_token=SecretStr(parsed.access_token), + expires_at_epoch=parsed.expires_at_epoch, + assertion_sha256=parsed.assertion_sha256, + ) + + def save(self, key: str, token: StoredToken) -> None: + if not self._usable(): + return + body: Final = ( + _StoredTokenFile( + access_token=token.access_token.get_secret_value(), + expires_at_epoch=token.expires_at_epoch, + assertion_sha256=token.assertion_sha256, + ) + .model_dump_json() + .encode() + ) + try: + _write_token_file(self._directory, key, body) + except OSError as e: + verbose_logger.debug("Token exchange cache entry not written: %s", e) + + def delete(self, key: str) -> None: + if not self._usable(): + return + with contextlib.suppress(FileNotFoundError, OSError): + (self._directory / f"{key}.json").unlink() + + @contextlib.contextmanager + def lock(self, key: str) -> Generator[None]: + if sys.platform == "win32" or not self._usable(): + yield + return + import fcntl + + try: + fd: Final = os.open(self._directory / f"{key}.lock", os.O_RDWR | os.O_CREAT, 0o600) + except OSError as e: + verbose_logger.debug("Token exchange cache lock unavailable (%s); minting without it", e) + yield + return + try: + fcntl.flock(fd, fcntl.LOCK_EX) + yield + finally: + with contextlib.suppress(OSError): + fcntl.flock(fd, fcntl.LOCK_UN) + os.close(fd) + + +def default_shared_token_store() -> SharedTokenStore | None: + """``LITELLM_TOKEN_EXCHANGE_CACHE_DIR`` relocates the store; setting it empty disables it. Without + it the store lives under the temp directory, keyed by uid, so the workers of one proxy share it and + other users on the host cannot read it. Windows has no ``flock``, so it caches per process there.""" + if sys.platform == "win32": + return None + configured: Final = os.environ.get(CACHE_DIR_ENV) + if configured == "": + return None + if configured is not None: + return FileTokenStore(Path(configured)) + return FileTokenStore(Path(tempfile.gettempdir()) / f"litellm-token-exchange-{os.getuid()}") diff --git a/litellm/llms/base_llm/auth/token_exchange.py b/litellm/llms/base_llm/auth/token_exchange.py new file mode 100644 index 00000000000..5f4b8c00bba --- /dev/null +++ b/litellm/llms/base_llm/auth/token_exchange.py @@ -0,0 +1,941 @@ +"""RFC 7523 JWT-bearer token exchange engine, shared across providers. + +One sync state machine per process: bounded engine-owned entry map, two-tier +refresh (advisory background refresh + mandatory single-flight), HTTPS pinning, +response caps, and RFC 6749 5.2 redaction. Providers describe a grant profile as +a ``TokenExchangeSpec`` and map the typed ``ExchangeError`` union to their own +public exception contract. +""" + +import asyncio +import hashlib +import json +import re +import threading +import time +from collections.abc import Callable, Coroutine, Iterator, Mapping, Sequence +from concurrent.futures import Executor, ThreadPoolExecutor +from dataclasses import dataclass +from itertools import chain +from math import inf +from types import MappingProxyType +from typing import TYPE_CHECKING, Final, Protocol, TypeAlias +from urllib.parse import unquote, unquote_plus, urlencode, urlsplit, urlunsplit + +import httpx +from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError +from typing_extensions import assert_never + +from litellm._logging import verbose_logger +from litellm.llms.base_llm.auth.shared_token_store import SharedTokenStore, StoredToken, default_shared_token_store +from litellm.llms.base_llm.auth.types import ( + AssertionReader, + AssertionSource, + AssertionSourceError, + ExchangeCallType, + ExchangeError, + ExchangeResult, + InsecureTokenUrl, + MalformedTokenResponse, + MintedToken, + SyncTokenPoster, + TokenEndpointError, + TokenExchangeMetricsSink, + TokenExchangeSpec, + TokenTransportError, +) +from litellm.types.services import ServiceTypes + +if TYPE_CHECKING: + from litellm.llms.custom_httpx.http_handler import HTTPHandler + +CALL_TYPE_COLD_MINT: Final[ExchangeCallType] = "cold_mint" +CALL_TYPE_MANDATORY_REFRESH: Final[ExchangeCallType] = "mandatory_refresh" +CALL_TYPE_ADVISORY_REFRESH: Final[ExchangeCallType] = "advisory_refresh" +CALL_TYPE_CACHE_HIT: Final = "cache_hit" + +ADVISORY_REFRESH_SECONDS: Final = 120.0 +MANDATORY_REFRESH_SECONDS: Final = 30.0 +ADVISORY_REFRESH_LIFETIME_FRACTION: Final = 0.5 +MANDATORY_REFRESH_LIFETIME_FRACTION: Final = 0.125 +ADVISORY_REFRESH_BACKOFF_SECONDS: Final = 5.0 +FALLBACK_TOKEN_TTL_SECONDS: Final = 60.0 +# Metrics are best-effort, so the backlog is capped and further events are dropped. Request volume +# must not be able to grow this queue without bound when a telemetry backend stalls. +_METRICS_QUEUE_LIMIT: Final = 1000 +MAX_ASSERTION_BYTES: Final = 16 * 1024 +MAX_RESPONSE_BYTES: Final = 1024 * 1024 + +_REDACTION_CAP: Final = 256 +_FOLLOWER_WAIT_GRACE_SECONDS: Final = 5.0 +_LOCAL_HOSTS: Final = frozenset({"localhost", "127.0.0.1", "::1"}) +_OAUTH_ERROR_FIELDS: Final = ("error", "error_description", "error_uri") +_NESTED_ERROR_FIELDS: Final = ("type", "message") +_CONTENT_TYPES: Final = MappingProxyType({"json": "application/json", "form": "application/x-www-form-urlencoded"}) +_OVERSIZED_BODY_MESSAGE: Final = "oversized error response omitted" +_NON_OBJECT_BODY_MESSAGE: Final = "non-object error response omitted" +_NO_OAUTH_FIELDS_MESSAGE: Final = "error response carried no RFC 6749 fields" +_UNSTRUCTURED_BODY_MESSAGE: Final = "non-JSON error response omitted" +_REFLECTED_VALUE_MESSAGE: Final = "" +# A credential fragment shorter than this is not worth the false positives; longer, and a run +# shared with the assertion is reflection rather than coincidence. +_REFLECTION_MIN_RUN: Final = 8 +# Everything a base64url credential is NOT made of, stripped so a fragment split by delimiters +# still lines up against the assertion. +_CREDENTIAL_CHARS: Final = re.compile(r"[^A-Za-z0-9._~+/=-]") +_SENTINEL_BODY_MESSAGES: Final = frozenset({_OVERSIZED_BODY_MESSAGE, _NON_OBJECT_BODY_MESSAGE}) + + +class _TokenExchangeResponse(BaseModel): + access_token: str + expires_in: int | None = None + token_type: str | None = None + + +_RedactableBody: TypeAlias = Mapping[str, object] | list[object] | str | int | float | bool | None +_REDACTABLE_BODY_ADAPTER: Final = TypeAdapter[_RedactableBody](_RedactableBody) + + +def endpoint_url_for_error_message(url: str) -> str: + """``url`` reduced to scheme, host and path for operator-facing errors. + + A token endpoint is configuration, not a secret, and naming it is what makes these errors + actionable. But nothing stops an operator writing a credential into it, as a query parameter + or as userinfo, and these errors reach model callers, so neither part is echoed. + """ + parsed: Final = urlsplit(url) + host: Final = parsed.hostname or "" + authority: Final = f"{host}:{parsed.port}" if parsed.port is not None else host + return urlunsplit((parsed.scheme, authority, parsed.path, "", "")) + + +def validate_token_endpoint_url(url: str) -> str | InsecureTokenUrl: + parsed: Final = urlsplit(url) + if parsed.scheme == "https": + return url + if parsed.scheme == "http" and (parsed.hostname or "") in _LOCAL_HOSTS: + return url + return InsecureTokenUrl(host=parsed.hostname or "") + + +def redact_oauth_error_body( + status_code: int, + body_text: str, + assertion: SecretStr | Sequence[SecretStr] | None = None, +) -> TokenEndpointError: + """``assertion`` may be every form of the credential that went out on the wire. + + A grant that encodes its credential before sending it (``client_secret_basic`` base64s + ``id:secret``) can have that encoded form echoed back, and it decodes straight to the secret, + so checking only the raw value lets reversible material through. + """ + rendered: Final = _redact_body_text(body_text) + secrets: Final = () if assertion is None else (assertion,) if isinstance(assertion, SecretStr) else tuple(assertion) + redacted: Final = next( + ( + _REFLECTED_VALUE_MESSAGE + for secret in secrets + if _drop_reflected_assertion(rendered, secret) is _REFLECTED_VALUE_MESSAGE + ), + rendered, + ) + return TokenEndpointError(status_code=status_code, redacted_body=redacted) + + +def _drop_reflected_assertion(rendered: str, assertion: SecretStr | None) -> str: + """Catches an endpoint that echoes the submitted credential back, verbatim or in fragments, + however it split or percent-encoded it. + + Both sides are reduced to the characters a credential is made of before comparison. Stripping + only the rendered side would stop matching a secret that carries spaces or punctuation of its + own, which is exactly the hand-set passphrase most at risk of being echoed. + + This stops an accidental or naive echo. It cannot stop an endpoint that deliberately re-encodes + or interleaves the credential, and it is not what keeps the credential from the endpoint, which + already holds it. What it protects is blast radius: keeping the value out of the caller's error + and out of third-party log sinks. + """ + if assertion is None: + return rendered + secret: Final = assertion.get_secret_value() + if not secret: + return rendered + if secret in rendered: + return _REFLECTED_VALUE_MESSAGE + compacted_secret: Final = _CREDENTIAL_CHARS.sub("", secret) + if not compacted_secret: + return rendered + return _REFLECTED_VALUE_MESSAGE if _shares_a_credential_run(rendered, compacted_secret) else rendered + + +def _shares_a_credential_run(rendered: str, compacted_secret: str) -> bool: + """``unquote`` covers a credential sent form-encoded, without every caller enumerating that + shape for itself: percent-escaping is reversible and applies to any field, query string + included. + + A secret shorter than the probe run is compared whole: a window longer than the secret can + never be found inside it, which would leave a short client secret unprotected in every shape + but the verbatim one. + """ + # unquote covers %XX; unquote_plus additionally covers the "+" a form-encoded body uses for a + # space. Both are kept rather than only the wider one, because "+" is a base64 character and + # decoding it away would lose a run that the undecoded candidate still matches on. + run: Final = min(_REFLECTION_MIN_RUN, len(compacted_secret)) + compacted_candidates: Final = tuple( + _CREDENTIAL_CHARS.sub("", candidate) for candidate in (rendered, unquote(rendered), unquote_plus(rendered)) + ) + windows: Final = chain.from_iterable(_character_runs(candidate, run) for candidate in compacted_candidates) + return any(window in compacted_secret for window in windows) + + +def _character_runs(compacted: str, run: int) -> Iterator[str]: + return (compacted[start : start + run] for start in range(len(compacted) - run + 1)) + + +def _redact_body_text(body_text: str) -> str: + if body_text in _SENTINEL_BODY_MESSAGES: + return body_text + if len(body_text) > MAX_RESPONSE_BYTES: + return _OVERSIZED_BODY_MESSAGE + try: + parsed: Final = _REDACTABLE_BODY_ADAPTER.validate_json(body_text) + except ValidationError: + return _UNSTRUCTURED_BODY_MESSAGE + match parsed: + case Mapping(): + return _format_oauth_error_fields(parsed) + case _: + return _NON_OBJECT_BODY_MESSAGE + + +def _format_oauth_error_fields(body: Mapping[str, object]) -> str: + fields: Final = tuple( + f"{name}: {_format_oauth_error_value(body[name])}" for name in _OAUTH_ERROR_FIELDS if body.get(name) is not None + ) + return "; ".join(fields) if fields else _NO_OAUTH_FIELDS_MESSAGE + + +def _format_oauth_error_value(value: object) -> str: + """RFC 6749 types ``error`` as a string, but Anthropic (and other providers) nest their + own ``{"type": ..., "message": ...}`` envelope there; render that rather than a dict repr.""" + if isinstance(value, Mapping): + nested: Final = tuple( + str(value[key])[:_REDACTION_CAP] for key in _NESTED_ERROR_FIELDS if value.get(key) is not None + ) + if nested: + return " - ".join(nested) + return str(value)[:_REDACTION_CAP] + + +def _error_summary(error: ExchangeError) -> str: + match error: + case AssertionSourceError(): + return f"AssertionSourceError: assertion {error.kind} from {error.source_ref}" + case InsecureTokenUrl(): + return f"InsecureTokenUrl: insecure token endpoint host {error.host}" + case TokenEndpointError(): + return f"TokenEndpointError: HTTP {error.status_code}: {error.redacted_body}" + case TokenTransportError(): + return f"TokenTransportError: {error.detail}" + case MalformedTokenResponse(): + return f"MalformedTokenResponse: {error.detail}" + case _: + assert_never(error) + + +class _MetricsFailure(Exception): + """Never raised: typed carriers handed to the service failure hook so the prometheus + ``error_class`` label names the ``ExchangeError`` variant; the message is the redacted + ``_error_summary`` and carries no credential material.""" + + +class TokenExchangeAssertionSourceFailure(_MetricsFailure): ... + + +class TokenExchangeInsecureUrlFailure(_MetricsFailure): ... + + +class TokenExchangeEndpointFailure(_MetricsFailure): ... + + +class TokenExchangeTransportFailure(_MetricsFailure): ... + + +class TokenExchangeMalformedResponseFailure(_MetricsFailure): ... + + +def _failure_exception(error: ExchangeError) -> _MetricsFailure: + summary: Final = _error_summary(error) + match error: + case AssertionSourceError(): + return TokenExchangeAssertionSourceFailure(summary) + case InsecureTokenUrl(): + return TokenExchangeInsecureUrlFailure(summary) + case TokenEndpointError(): + return TokenExchangeEndpointFailure(summary) + case TokenTransportError(): + return TokenExchangeTransportFailure(summary) + case MalformedTokenResponse(): + return TokenExchangeMalformedResponseFailure(summary) + case _: + assert_never(error) + + +def _cache_key(spec: TokenExchangeSpec) -> str: + return hashlib.sha256( + "\x1f".join((spec.token_url, spec.assertion_ref, *spec.cache_key_identity)).encode() + ).hexdigest() + + +def _shares_one_assertion_across_workers(spec: TokenExchangeSpec) -> bool: + """The store exists so the workers reading one projected token file don't each spend that file's + single-use ``jti``. A source that mints its own assertion per exchange shares nothing with another + worker, so it never reads the store, never finds a hit there, and keeps its minted token off disk.""" + return spec.assertion_source is None + + +def _assertion_digest(assertion: SecretStr) -> str: + return hashlib.sha256(assertion.get_secret_value().encode()).hexdigest() + + +def _assertion_fetch(reader: AssertionReader, spec: TokenExchangeSpec) -> AssertionSource: + """``spec.assertion_source`` (an identity source's own fetch/mint closure) takes priority over + the engine-level reader when set; either way, failures are reported against ``spec.assertion_ref``.""" + if spec.assertion_source is not None: + return spec.assertion_source + return lambda: reader(spec.assertion_ref) + + +def _read_assertion(fetch: AssertionSource, ref: str) -> SecretStr | AssertionSourceError: + from litellm.secret_managers.main import OidcPathNotAllowedError + + try: + raw: Final = fetch() + except OidcPathNotAllowedError: + return AssertionSourceError(kind="disallowed_path", source_ref=ref) + except (ValueError, ImportError) as e: + return AssertionSourceError(kind="unreadable", source_ref=ref, detail=str(e)[:_REDACTION_CAP]) + except Exception: # noqa: BLE001 # injected readers (secret managers) raise arbitrarily; all failures become values + return AssertionSourceError(kind="unreadable", source_ref=ref) + if raw is None: + return AssertionSourceError(kind="missing", source_ref=ref) + stripped: Final = raw.strip() + if not stripped: + return AssertionSourceError(kind="empty", source_ref=ref) + if len(stripped.encode("utf-8")) > MAX_ASSERTION_BYTES: + return AssertionSourceError(kind="oversized", source_ref=ref) + return SecretStr(stripped) + + +def _serialize_body(spec: TokenExchangeSpec, assertion: SecretStr) -> bytes: + if spec.body_encoding == "json": + return json.dumps( + { + **spec.static_body, + spec.assertion_field: assertion.get_secret_value(), + } + ).encode() + return urlencode( + { + **spec.static_body, + spec.assertion_field: assertion.get_secret_value(), + } + ).encode() + + +def _sanitize_expires_in(expires_in: int | None) -> float: + if expires_in is None or expires_in <= 0: + return FALLBACK_TOKEN_TTL_SECONDS + return float(expires_in) + + +@dataclass(frozen=True, slots=True) +class _RefreshWindows: + advisory: float + mandatory: float + + +def _refresh_windows(lifetime_seconds: float | None) -> _RefreshWindows: + """A token whose whole life is shorter than the flat windows sits inside them from the moment it + is minted, so every request would arm another background exchange against the token endpoint. + Scaling each window by a fraction of the observed lifetime makes a 60s token refresh around its + half life instead; at a lifetime of 240s and above both fractions reach the flat windows, so + ordinary long-lived tokens keep exactly the 120s/30s behaviour.""" + if lifetime_seconds is None or lifetime_seconds <= 0.0: + return _RefreshWindows(advisory=ADVISORY_REFRESH_SECONDS, mandatory=MANDATORY_REFRESH_SECONDS) + return _RefreshWindows( + advisory=min(ADVISORY_REFRESH_SECONDS, lifetime_seconds * ADVISORY_REFRESH_LIFETIME_FRACTION), + mandatory=min(MANDATORY_REFRESH_SECONDS, lifetime_seconds * MANDATORY_REFRESH_LIFETIME_FRACTION), + ) + + +def _capped_body_text(response: httpx.Response) -> str: + if len(response.content) > MAX_RESPONSE_BYTES: + return _OVERSIZED_BODY_MESSAGE + return response.text + + +def _default_assertion_reader(ref: str) -> str | None: + from litellm.secret_managers.main import get_secret_str + + return get_secret_str(ref) + + +def _new_exchange_handler() -> "HTTPHandler": + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + return HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0), follow_redirects=False) + + +def require_posted_response(response: httpx.Response | None, endpoint_label: str) -> httpx.Response: + """The legacy ``HTTPHandler`` carries no return annotation, so a patched or stubbed client can + hand a poster ``None`` back; a transport error beats dereferencing it.""" + if response is None: + raise httpx.TransportError(f"{endpoint_label} returned no response") + return response + + +class _HttpxSyncTokenPoster: + """Default poster: a dedicated HTTPHandler (no logging_obj, so litellm's + pre/post-call body logging never sees the exchange POST); returns the + response for any status.""" + + def __init__(self, handler_factory: Callable[[], "HTTPHandler"] = _new_exchange_handler) -> None: + self._lock: Final = threading.Lock() + self._handler_factory: Final = handler_factory + self._handler: HTTPHandler | None = None + + def _handler_instance(self) -> "HTTPHandler": + with self._lock: + if self._handler is None: + self._handler = self._handler_factory() + return self._handler + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + try: + response: Final[httpx.Response | None] = self._handler_instance().post( # pyright: ignore[reportUnknownMemberType] # HTTPHandler.post is legacy-untyped; the result is validated below + url, + content=content, + headers=dict(headers), + timeout=timeout, + ) + except httpx.HTTPStatusError as e: + return e.response + return require_posted_response(response, "token endpoint") + + +class _ServiceLoggingHooks(Protocol): + """The slice of ``litellm._service_logger.ServiceLogging`` the metrics sink calls; a protocol + so tests inject a recorder instead of monkeypatching.""" + + async def async_service_success_hook(self, service: ServiceTypes, call_type: str, duration: float) -> None: ... + + async def async_service_failure_hook( + self, service: ServiceTypes, duration: float, error: str | Exception, call_type: str + ) -> None: ... + + +_HooksCoroFactory: TypeAlias = Callable[ + [_ServiceLoggingHooks], + Coroutine[object, object, None], +] + + +def _default_service_logging() -> _ServiceLoggingHooks: + from litellm._service_logger import ServiceLogging + + return ServiceLogging() + + +class ServiceLoggingMetricsSink: + """Default sink: bridges engine metrics onto litellm's ServiceTypes pattern + (prometheus ``litellm_anthropic_wif_*`` via ``service_callback``). The engine's entry points + are sync threads with no event loop, and the service hooks are async, so every emission is + fire-and-forget on a dedicated single worker thread that owns its own short-lived loop -- + the mint path only ever pays for an executor queue put.""" + + def __init__( + self, + service_logging_factory: Callable[[], _ServiceLoggingHooks] = _default_service_logging, + executor: Executor | None = None, + ) -> None: + self._lock: Final = threading.Lock() + self._service_logging_factory: Final = service_logging_factory + self._service_logging: _ServiceLoggingHooks | None = None + self._executor: Executor | None = executor + self._queued: int = 0 + + def _service_logging_instance(self) -> _ServiceLoggingHooks: + with self._lock: + if self._service_logging is None: + self._service_logging = self._service_logging_factory() + return self._service_logging + + def _executor_instance(self) -> Executor: + with self._lock: + if self._executor is None: + self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="litellm-token-exchange-metrics") + return self._executor + + def _emit(self, coro_factory: _HooksCoroFactory) -> None: + try: + asyncio.run(coro_factory(self._service_logging_instance())) + except Exception as e: # noqa: BLE001 # metrics are best-effort; emission failures must never surface + verbose_logger.debug("token exchange metrics emission failed: %s", e) + + def _submit(self, coro_factory: _HooksCoroFactory) -> None: + """Drop the event rather than queue it once the backlog is full. A stalled telemetry + backend must not let request volume grow an unbounded queue in the proxy: losing a + metric sample is always cheaper than losing the process.""" + with self._lock: + if self._queued >= _METRICS_QUEUE_LIMIT: + verbose_logger.debug("token exchange metrics queue full, dropping event") + return + self._queued += 1 + try: + self._executor_instance().submit(self._emit_and_release, coro_factory) + except Exception as e: # noqa: BLE001 # a rejected submit must not surface to the mint + with self._lock: + self._queued -= 1 + verbose_logger.debug("token exchange metrics submit failed: %s", e) + + def _emit_and_release(self, coro_factory: _HooksCoroFactory) -> None: + try: + self._emit(coro_factory) + finally: + with self._lock: + self._queued -= 1 + + def exchange_success(self, *, call_type: ExchangeCallType, duration_seconds: float) -> None: + def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]: + return hooks.async_service_success_hook( + service=ServiceTypes.ANTHROPIC_WIF, call_type=call_type, duration=duration_seconds + ) + + self._submit(start) + + def exchange_failure(self, *, call_type: ExchangeCallType, duration_seconds: float, error: ExchangeError) -> None: + failure: Final = _failure_exception(error) + + def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]: + return hooks.async_service_failure_hook( + service=ServiceTypes.ANTHROPIC_WIF, duration=duration_seconds, error=failure, call_type=call_type + ) + + self._submit(start) + + def cache_hit(self) -> None: + def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]: + return hooks.async_service_success_hook( + service=ServiceTypes.ANTHROPIC_WIF_CACHE, call_type=CALL_TYPE_CACHE_HIT, duration=0.0 + ) + + self._submit(start) + + +class _Entry: + """Single-flight state for one cache key; mutable by design, confined to the + engine, and only ever mutated under the engine lock.""" + + __slots__ = ("backoff_until", "done", "force_refresh", "in_flight", "last_error", "lifetime_seconds", "token") + + def __init__(self, force_refresh: bool = False) -> None: + self.token: MintedToken | None = None + self.lifetime_seconds: float | None = None + self.in_flight: bool = False + self.done: Final = threading.Event() + self.backoff_until: float = float("-inf") + self.force_refresh: bool = force_refresh + self.last_error: ExchangeError | None = None + + def arm(self) -> None: + self.in_flight = True + self.last_error = None + self.done.clear() + + def disarm(self, now: float) -> None: + """Undo ``arm`` for a refresh that never started. Nothing is on its way to publish, so the + entry must stop reading as in-flight, and the backoff keeps every later caller from + re-attempting a schedule that just failed.""" + self.backoff_until = now + ADVISORY_REFRESH_BACKOFF_SECONDS + self.in_flight = False + self.done.set() + + def _store(self, token: MintedToken, now: float) -> None: + self.token = token + self.lifetime_seconds = None if token.expires_at is None else max(token.expires_at - now, 0.0) + self.last_error = None + + def publish(self, result: ExchangeResult, now: float) -> None: + match result: + case MintedToken(): + self._store(result, now) + case _: + self.last_error = result + self.backoff_until = now + ADVISORY_REFRESH_BACKOFF_SECONDS + self.force_refresh = False + self.in_flight = False + self.done.set() + + def publish_advisory(self, result: ExchangeResult, now: float) -> None: + """A failed advisory refresh records only the backoff, never ``last_error``: a follower whose + cached token expires while this runs must be free to re-lead a fresh mint and recover.""" + match result: + case MintedToken(): + self._store(result, now) + case _: + self.backoff_until = now + ADVISORY_REFRESH_BACKOFF_SECONDS + self.in_flight = False + self.done.set() + + +@dataclass(frozen=True, slots=True) +class _Serve: + token: MintedToken + + +@dataclass(frozen=True, slots=True) +class _ServeAndRefresh: + token: MintedToken + + +@dataclass(frozen=True, slots=True) +class _Lead: + call_type: ExchangeCallType + + +@dataclass(frozen=True, slots=True) +class _Follow: + pass + + +@dataclass(frozen=True, slots=True) +class _Fail: + error: ExchangeError + + +_Decision: TypeAlias = _Serve | _ServeAndRefresh | _Lead | _Follow | _Fail + + +@dataclass(frozen=True, slots=True) +class _Unauthorized: + response: httpx.Response + assertion: SecretStr + + +def _denied(attempt: _Unauthorized) -> TokenEndpointError: + return redact_oauth_error_body(attempt.response.status_code, _capped_body_text(attempt.response), attempt.assertion) + + +class JwtBearerTokenExchangeEngine: + def __init__( + self, + poster: SyncTokenPoster | None = None, + assertion_reader: AssertionReader | None = None, + clock: Callable[[], float] = time.monotonic, + refresh_executor: Executor | None = None, + max_entries: int = 64, + metrics_sink: TokenExchangeMetricsSink | None = None, + shared_store: SharedTokenStore | None = None, + wall_clock: Callable[[], float] = time.time, + ) -> None: + self._poster: Final[SyncTokenPoster] = poster if poster is not None else _HttpxSyncTokenPoster() + self._assertion_reader: Final[AssertionReader] = ( + assertion_reader if assertion_reader is not None else _default_assertion_reader + ) + self._clock: Final = clock + self._refresh_executor: Executor | None = refresh_executor + self._max_entries: Final = max_entries + self._metrics_sink: Final[TokenExchangeMetricsSink] = ( + metrics_sink if metrics_sink is not None else ServiceLoggingMetricsSink() + ) + self._shared_store: Final = shared_store + self._wall_clock: Final = wall_clock + self._lock: Final = threading.Lock() + self._entries: Final[dict[str, _Entry]] = {} # mutable-ok: engine-owned map guarded by _lock + + def get_token(self, spec: TokenExchangeSpec) -> ExchangeResult: + """A follower whose leader published nothing re-classifies rather than recursing, so a + contended entry cannot grow the stack one frame per failed leader.""" + while True: + with self._lock: + entry = self._get_or_create_entry_locked(spec) + decision = self._classify_and_arm_locked(entry) + match decision: + case _Serve(token=token): + self._report_cache_hit() + return token + case _ServeAndRefresh(token=token): + self._report_cache_hit() + self._submit_advisory_refresh(spec, entry) + return token + case _Fail(error=error): + return error + case _Lead(call_type=call_type): + return self._lead(spec, entry, call_type) + case _Follow(): + followed = self._await_leader(spec, entry) + if followed is not None: + return followed + case _: + assert_never(decision) + + async def aget_token(self, spec: TokenExchangeSpec) -> ExchangeResult: + return await asyncio.to_thread(self.get_token, spec) + + def invalidate(self, spec: TokenExchangeSpec) -> None: + key: Final = _cache_key(spec) + with self._lock: + if key in self._entries: + self._entries[key] = _Entry(force_refresh=True) + if self._shared_store is not None: + self._shared_store.delete(key) + + def _get_or_create_entry_locked(self, spec: TokenExchangeSpec) -> _Entry: + key: Final = _cache_key(spec) + existing: Final = self._entries.get(key) + if existing is not None: + return existing + if len(self._entries) >= self._max_entries: + self._evict_locked() + created: Final = _Entry() + self._entries[key] = created + return created + + def _evict_locked(self) -> None: + now: Final = self._clock() + stale: Final = tuple( + key + for key, entry in self._entries.items() + if not entry.in_flight + and (entry.token is None or (entry.token.expires_at is not None and entry.token.expires_at <= now)) + ) + for key in stale: + del self._entries[key] + if len(self._entries) < self._max_entries: + return + # Evict soonest-to-expire first, and take as many as the overshoot needs rather than one, so a + # burst of distinct identities does not leave the map permanently above max_entries. An entry + # a leader owns or a follower waits on is never a candidate, so a moment where every entry is + # in flight still over-inserts; that residue is bounded by the concurrent mints themselves. + evictable: Final = sorted( + ( + entry.token.expires_at if entry.token is not None and entry.token.expires_at is not None else -inf, + key, + ) + for key, entry in self._entries.items() + if not entry.in_flight + ) + for _, key in evictable[: len(self._entries) - self._max_entries + 1]: + del self._entries[key] + + def _classify_and_arm_locked(self, entry: _Entry) -> _Decision: + token: Final = entry.token + if token is not None and not entry.force_refresh: + if token.expires_at is None: + return _Serve(token=token) + windows: Final = _refresh_windows(entry.lifetime_seconds) + remaining: Final = token.expires_at - self._clock() + if remaining > windows.advisory: + return _Serve(token=token) + if remaining > windows.mandatory: + if entry.in_flight or self._clock() < entry.backoff_until: + return _Serve(token=token) + entry.arm() + return _ServeAndRefresh(token=token) + if entry.in_flight: + return _Follow() + if entry.last_error is not None and self._clock() < entry.backoff_until: + return _Fail(error=entry.last_error) + entry.arm() + return _Lead(call_type=CALL_TYPE_COLD_MINT if token is None else CALL_TYPE_MANDATORY_REFRESH) + + def _executor_instance(self) -> Executor: + with self._lock: + if self._refresh_executor is None: + self._refresh_executor = ThreadPoolExecutor(thread_name_prefix="litellm-token-exchange-refresh") + return self._refresh_executor + + def _submit_advisory_refresh(self, spec: TokenExchangeSpec, entry: _Entry) -> None: + """The entry is already armed, so an executor that refuses the work would leave it reading + as in-flight with nothing on its way to publish, and every later caller would wait out the + follower timeout and fail. A refused submit disarms it and the cached token keeps serving.""" + try: + self._executor_instance().submit(self._advisory_refresh, spec, entry) + except RuntimeError as e: + verbose_logger.debug("token exchange advisory refresh could not be scheduled: %s", e) + with self._lock: + entry.disarm(self._clock()) + + def _lead(self, spec: TokenExchangeSpec, entry: _Entry, call_type: ExchangeCallType) -> ExchangeResult: + started: Final = self._clock() + result: Final = self._exchange_never_raises(spec) + duration: Final = self._clock() - started + with self._lock: + entry.publish(result, now=self._clock()) + self._report_exchange(call_type, duration, result) + return result + + def _await_leader(self, spec: TokenExchangeSpec, entry: _Entry) -> "ExchangeResult | None": + """None means the finished round left neither a valid token nor an error + (a failed advisory refresh); the caller re-enters and leads a fresh exchange.""" + leader_finished: Final = entry.done.wait(2 * spec.timeout_seconds + _FOLLOWER_WAIT_GRACE_SECONDS) + with self._lock: + token: Final = entry.token + if token is not None and (token.expires_at is None or token.expires_at > self._clock()): + return token + if entry.last_error is not None: + return entry.last_error + if leader_finished: + return None + return TokenTransportError(detail="timed out waiting for the token exchange leader") + + def _advisory_refresh(self, spec: TokenExchangeSpec, entry: _Entry) -> None: + started: Final = self._clock() + result: Final = self._exchange_never_raises(spec) + duration: Final = self._clock() - started + with self._lock: + now: Final = self._clock() + entry.publish_advisory(result, now=now) + stale_expires_at: Final = entry.token.expires_at if entry.token is not None else None + stale_mandatory: Final = _refresh_windows(entry.lifetime_seconds).mandatory + self._report_exchange(CALL_TYPE_ADVISORY_REFRESH, duration, result) + if isinstance(result, MintedToken): + return + seconds_to_mandatory_wall: Final = ( + max(stale_expires_at - now - stale_mandatory, 0.0) if stale_expires_at is not None else 0.0 + ) + verbose_logger.warning( + "Advisory token refresh against %s failed (%s); serving the cached token for up to " + "%.0fs before the mandatory refresh wall; next attempt after %.0fs backoff", + urlsplit(spec.token_url).hostname or "", + _error_summary(result), + seconds_to_mandatory_wall, + ADVISORY_REFRESH_BACKOFF_SECONDS, + ) + + def _report_exchange(self, call_type: ExchangeCallType, duration_seconds: float, result: ExchangeResult) -> None: + try: + match result: + case MintedToken(): + self._metrics_sink.exchange_success(call_type=call_type, duration_seconds=duration_seconds) + case _: + self._metrics_sink.exchange_failure( + call_type=call_type, duration_seconds=duration_seconds, error=result + ) + except Exception as e: # noqa: BLE001 # metrics are best-effort; a sink failure must never fail a mint + verbose_logger.debug("token exchange metrics emission failed: %s", e) + + def _report_cache_hit(self) -> None: + try: + self._metrics_sink.cache_hit() + except Exception as e: # noqa: BLE001 # metrics are best-effort; a sink failure must never fail a serve + verbose_logger.debug("token exchange cache-hit metric emission failed: %s", e) + + def _exchange_never_raises(self, spec: TokenExchangeSpec) -> ExchangeResult: + """The single-flight leader and the advisory refresher must always publish a result: an + unhandled exception here would leave the entry armed (in_flight, cleared event) forever, so + every subsequent caller for this key would follow a leader that never finishes.""" + try: + return self._exchange(spec) + except Exception as e: # noqa: BLE001 # a leader must resolve its entry; any failure becomes a value + return TokenTransportError(detail=f"{type(e).__name__}: {e}"[:_REDACTION_CAP]) + + def _exchange(self, spec: TokenExchangeSpec) -> ExchangeResult: + url_check: Final = validate_token_endpoint_url(spec.token_url) + if isinstance(url_check, InsecureTokenUrl): + return url_check + fetch: Final = _assertion_fetch(self._assertion_reader, spec) + assertion: Final = _read_assertion(fetch, spec.assertion_ref) + if isinstance(assertion, AssertionSourceError): + return assertion + if self._shared_store is None or not _shares_one_assertion_across_workers(spec): + return self._mint(spec, fetch, assertion) + key: Final = _cache_key(spec) + with self._shared_store.lock(key): + shared: Final = self._shared_token(self._shared_store.load(key), _assertion_digest(assertion)) + if shared is not None: + return shared + minted: Final = self._mint(spec, fetch, assertion) + if isinstance(minted, MintedToken): + self._shared_store.save(key, self._stored_token(minted)) + return minted + + def _shared_token(self, stored: StoredToken | None, assertion_sha256: str) -> MintedToken | None: + """A stored token minted from the very assertion this process holds is the token that assertion + bought: another worker sharing the token file already exchanged it, and an issuer enforcing + single-use ``jti`` would only deny a second exchange.""" + if stored is None or stored.assertion_sha256 != assertion_sha256: + return None + if stored.expires_at_epoch is None: + return MintedToken(access_token=stored.access_token, expires_at=None, assertion_sha256=assertion_sha256) + remaining: Final = stored.expires_at_epoch - self._wall_clock() + if remaining <= 0.0: + return None + return MintedToken( + access_token=stored.access_token, + expires_at=self._clock() + remaining, + assertion_sha256=assertion_sha256, + ) + + def _stored_token(self, token: MintedToken) -> StoredToken: + return StoredToken( + access_token=token.access_token, + expires_at_epoch=( + None if token.expires_at is None else self._wall_clock() + (token.expires_at - self._clock()) + ), + assertion_sha256=token.assertion_sha256, + ) + + def _mint(self, spec: TokenExchangeSpec, fetch: AssertionSource, assertion: SecretStr) -> ExchangeResult: + """One 401 earns one retry, and only with an assertion that changed since the first attempt: a + token file rotated between the read and the POST is worth resending, the same assertion is not, + since an issuer that already consumed its ``jti`` denies it again.""" + first: Final = self._post_assertion(spec, assertion) + if not isinstance(first, _Unauthorized): + return first + reread: Final = _read_assertion(fetch, spec.assertion_ref) + if isinstance(reread, AssertionSourceError): + return reread + if reread.get_secret_value() == assertion.get_secret_value(): + return _denied(first) + second: Final = self._post_assertion(spec, reread) + if isinstance(second, _Unauthorized): + return _denied(second) + return second + + def _post_assertion(self, spec: TokenExchangeSpec, assertion: SecretStr) -> "ExchangeResult | _Unauthorized": + try: + response: Final = self._poster.post( + spec.token_url, + content=_serialize_body(spec, assertion), + headers=MappingProxyType({"content-type": _CONTENT_TYPES[spec.body_encoding], **spec.request_headers}), + timeout=spec.timeout_seconds, + ) + except Exception as e: # noqa: BLE001 # injected posters may raise beyond httpx; transport failures become values + return TokenTransportError(detail=f"{type(e).__name__}: {e}"[:_REDACTION_CAP]) + if response.status_code == 401: + return _Unauthorized(response=response, assertion=assertion) + return self._parse_response(response, assertion) + + def _parse_response(self, response: httpx.Response, assertion: SecretStr) -> ExchangeResult: + if not 200 <= response.status_code < 300: + return redact_oauth_error_body(response.status_code, _capped_body_text(response), assertion) + if len(response.content) > MAX_RESPONSE_BYTES: + return MalformedTokenResponse(detail="token response body exceeds the 1 MiB cap") + try: + parsed: Final = _TokenExchangeResponse.model_validate_json(response.content) + except ValidationError: + return MalformedTokenResponse(detail="token response failed RFC 6749 5.1 schema validation") + if parsed.token_type is not None and parsed.token_type.lower() != "bearer": + return MalformedTokenResponse(detail="token response carried a non-bearer token_type") + if not parsed.access_token.strip(): + return MalformedTokenResponse(detail="token response carried an empty access_token") + return MintedToken( + access_token=SecretStr(parsed.access_token), + expires_at=self._clock() + _sanitize_expires_in(parsed.expires_in), + assertion_sha256=_assertion_digest(assertion), + ) + + +default_token_exchange_engine: Final = JwtBearerTokenExchangeEngine(shared_store=default_shared_token_store()) diff --git a/litellm/llms/base_llm/auth/types.py b/litellm/llms/base_llm/auth/types.py new file mode 100644 index 00000000000..c074ce57a77 --- /dev/null +++ b/litellm/llms/base_llm/auth/types.py @@ -0,0 +1,100 @@ +"""Provider-agnostic types for the RFC 7523 JWT-bearer token exchange engine.""" + +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from typing import Literal, Protocol, TypeAlias + +import httpx +from pydantic import SecretStr + +BodyEncoding: TypeAlias = Literal["json", "form"] +AssertionReader: TypeAlias = Callable[[str], str | None] +AssertionSource: TypeAlias = Callable[[], str | None] + + +@dataclass(frozen=True, slots=True) +class TokenExchangeSpec: + """One grant profile as pure data: one instance per (provider, deployment, identity). + + ``token_url`` must be derived from deployment config/env only, never per-request caller + input. ``assertion_ref`` is a ``oidc/...`` get_secret ref resolved fresh on every exchange. + + ``assertion_source``, when set, is a zero-arg per-config fetch/mint closure that the engine + prefers over its own engine-level ``AssertionReader`` -- the dispatch mechanism identity + sources beyond token_file/env (e.g. ``internal_issuer``, ``keycloak``) use to plug into the + shared engine without a global registry. ``assertion_ref`` still names the cache-key + discriminator and the ref echoed into operator-facing errors either way. + """ + + token_url: str + assertion_ref: str + assertion_field: str + static_body: Mapping[str, str] + body_encoding: BodyEncoding + request_headers: Mapping[str, str] + cache_key_identity: tuple[str, ...] + timeout_seconds: float = 30.0 + assertion_source: AssertionSource | None = None + + +@dataclass(frozen=True, slots=True) +class MintedToken: + access_token: SecretStr + expires_at: float | None + assertion_sha256: str + + +@dataclass(frozen=True, slots=True) +class AssertionSourceError: + kind: Literal["missing", "empty", "oversized", "unreadable", "disallowed_path"] + source_ref: str + detail: str | None = None + + +@dataclass(frozen=True, slots=True) +class InsecureTokenUrl: + host: str + + +@dataclass(frozen=True, slots=True) +class TokenEndpointError: + status_code: int + redacted_body: str + + +@dataclass(frozen=True, slots=True) +class TokenTransportError: + detail: str + + +@dataclass(frozen=True, slots=True) +class MalformedTokenResponse: + detail: str + + +ExchangeError: TypeAlias = ( + AssertionSourceError | InsecureTokenUrl | TokenEndpointError | TokenTransportError | MalformedTokenResponse +) +ExchangeResult: TypeAlias = MintedToken | ExchangeError + +ExchangeCallType: TypeAlias = Literal["cold_mint", "mandatory_refresh", "advisory_refresh"] + + +class TokenExchangeMetricsSink(Protocol): + """Observability seam for the exchange engine. Implementations must be best-effort: never raise + into the mint path, never block the calling thread, and never receive credential material -- + ``ExchangeError`` values are redacted by construction.""" + + def exchange_success(self, *, call_type: ExchangeCallType, duration_seconds: float) -> None: ... + + def exchange_failure( + self, *, call_type: ExchangeCallType, duration_seconds: float, error: ExchangeError + ) -> None: ... + + def cache_hit(self) -> None: ... + + +class SyncTokenPoster(Protocol): + """Returns the response for ANY status; never raises for status.""" + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: ... diff --git a/litellm/llms/base_llm/base_utils.py b/litellm/llms/base_llm/base_utils.py index a446a721307..9f690b9825a 100644 --- a/litellm/llms/base_llm/base_utils.py +++ b/litellm/llms/base_llm/base_utils.py @@ -5,6 +5,7 @@ Utility functions for base LLM classes. import copy import json from abc import ABC, abstractmethod +from collections.abc import Mapping from typing import Any, Final from openai.lib import _parsing, _pydantic @@ -65,6 +66,22 @@ class BaseLLMModelInfo(ABC): """ return [] + def discover_models( + self, litellm_params: Mapping[str, object] | None = None + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + """ + Live model discovery for a configured deployment. Defaults to the api_key/api_base + facade every provider already implements via ``get_models``; a provider whose + discovery needs more of ``litellm_params`` (e.g. Anthropic's workload identity + federation) overrides this instead of widening ``get_models`` for every provider. + """ + api_key: Final = litellm_params.get("api_key") if litellm_params is not None else None + api_base: Final = litellm_params.get("api_base") if litellm_params is not None else None + return self.get_models( + api_key=api_key if isinstance(api_key, str) else None, + api_base=api_base if isinstance(api_base, str) else None, + ) + @staticmethod @abstractmethod def get_api_key(api_key: str | None = None) -> str | None: diff --git a/litellm/llms/custom_httpx/http_handler.py b/litellm/llms/custom_httpx/http_handler.py index 312f46bb021..8235a116d70 100644 --- a/litellm/llms/custom_httpx/http_handler.py +++ b/litellm/llms/custom_httpx/http_handler.py @@ -1382,10 +1382,12 @@ class HTTPHandler: ssl_verify: bool | str | None = None, disable_default_headers: bool | None = False, # arize phoenix returns different API responses when user agent header in request + follow_redirects: bool = True, ): self.timeout = timeout self.ssl_verify = ssl_verify self.disable_default_headers = disable_default_headers + self.follow_redirects = follow_redirects self._owns_client = client is None self._heal_lock = threading.Lock() self._client = self.create_client() if client is None else client @@ -1410,7 +1412,7 @@ class HTTPHandler: cert=cert, headers=default_headers, cookies=blocked_cookie_jar(), - follow_redirects=True, + follow_redirects=self.follow_redirects, http2=http2_enabled(), ) @@ -1436,7 +1438,7 @@ class HTTPHandler: self, url: str, params: dict | None = None, - headers: dict | None = None, + headers: Mapping[str, Any] | None = None, follow_redirects: bool | None = None, timeout: float | httpx.Timeout | None = None, ): diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index 72bbb4f9556..7dd3b9b34c0 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1,4 +1,5 @@ import asyncio +import inspect import json import ssl from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Coroutine, Iterator, Mapping, Sequence @@ -18,6 +19,7 @@ from typing import ( Union, cast, get_type_hints, + runtime_checkable, ) from urllib.parse import parse_qs, urlencode, urlparse, urlunparse @@ -277,6 +279,55 @@ class _MediaUploadKwargs(TypedDict, total=False): timeout: float | httpx.Timeout +@runtime_checkable +class _AsyncFilesEnvironmentValidator(Protocol): + async def avalidate_environment( + self, + headers: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + model: str, + messages: list, # mutable-ok: mirrors the sync validate_environment contract this overrides + optional_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + litellm_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: ... # mutable-ok: mirrors the sync validate_environment contract this overrides + + +async def _avalidate_files_environment( + provider_config: BaseFilesConfig | BaseBatchesConfig, + *, + headers: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + model: str, + messages: list, # mutable-ok: mirrors the sync validate_environment contract this overrides + optional_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + litellm_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides + api_key: str | None, +) -> dict: # mutable-ok: mirrors the sync validate_environment contract this overrides + """Await the provider's async credential hook when it has one (e.g. Anthropic's workload + identity token exchange); otherwise offload the sync hook to a worker thread. Either way + the caller, an async file handler, never blocks the event loop on it.""" + if isinstance(provider_config, _AsyncFilesEnvironmentValidator) and inspect.iscoroutinefunction( + provider_config.avalidate_environment + ): + return await provider_config.avalidate_environment( + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=api_key, + ) + return await asyncio.to_thread( + provider_config.validate_environment, + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=api_key, + ) + + class _SignedBodyKwargs(TypedDict, total=False): data: ReadOnly[bytes] json: ReadOnly[dict[str, object]] @@ -387,6 +438,21 @@ class _PreparedFileContentRequest(NamedTuple): headers: dict +def _logged_file_content_request( + url: str, + params: dict, + request_headers: dict, + file_content_request: "FileContentRequest", + logging_obj: LiteLLMLoggingObj, +) -> _PreparedFileContentRequest: + logging_obj.pre_call( + input="", + api_key="", + additional_args={"api_base": url, "headers": request_headers, "file_id": file_content_request.get("file_id")}, + ) + return _PreparedFileContentRequest(url=url, params=params, headers=request_headers) + + async def _aiter_bytes_then_close(response: httpx.Response, *, chunk_size: int) -> AsyncGenerator[bytes, None]: try: async for chunk in response.aiter_bytes(chunk_size=chunk_size): @@ -2027,7 +2093,7 @@ class BaseLLMHTTPHandler: ( headers, api_base, - ) = anthropic_messages_provider_config.validate_anthropic_messages_environment( + ) = await anthropic_messages_provider_config.avalidate_anthropic_messages_environment( headers=merged_headers or {}, model=model, messages=messages, @@ -3324,6 +3390,19 @@ class BaseLLMHTTPHandler: """ Creates a file using Gemini's two-step upload process """ + if _is_async: + return self._avalidate_and_create_file( + create_file_data=create_file_data, + litellm_params=litellm_params, + provider_config=provider_config, + headers=headers, + api_base=api_base, + api_key=api_key, + logging_obj=logging_obj, + client=client, + timeout=timeout, + ) + # get config from model, custom llm provider headers = provider_config.validate_environment( api_key=api_key, @@ -3353,18 +3432,6 @@ class BaseLLMHTTPHandler: optional_params={}, ) - if _is_async: - return self.async_create_file( - transformed_request=transformed_request, - litellm_params=litellm_params, - provider_config=provider_config, - headers=headers, - api_base=api_base, - logging_obj=logging_obj, - client=client, - timeout=timeout, - ) - if client is None or not isinstance(client, HTTPHandler): sync_httpx_client = _get_httpx_client() else: @@ -3492,6 +3559,54 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params_with_url, ) + async def _avalidate_and_create_file( + self, + *, + create_file_data: CreateFileRequest, + litellm_params: dict, # mutable-ok: mirrors the create_file contract this dispatches for + provider_config: BaseFilesConfig, + headers: dict, # mutable-ok: mirrors the create_file contract this dispatches for + api_base: str | None, + api_key: str | None, + logging_obj: LiteLLMLoggingObj, + client: HTTPHandler | AsyncHTTPHandler | None, + timeout: float | httpx.Timeout | None, + ) -> OpenAIFileObject: + validated_headers: Final = await _avalidate_files_environment( + provider_config, + headers=headers, + model="", + messages=[], + optional_params={}, + litellm_params=litellm_params, + api_key=api_key, + ) + complete_api_base: Final = provider_config.get_complete_file_url( + api_base=api_base, + api_key=api_key, + model="", + optional_params={}, + litellm_params=litellm_params, + data=create_file_data, + ) + if not complete_api_base: + raise ValueError("api_base is required for create_file") + return await self.async_create_file( + transformed_request=provider_config.transform_create_file_request( + model="", + create_file_data=create_file_data, + litellm_params=litellm_params, + optional_params={}, + ), + litellm_params=litellm_params, + provider_config=provider_config, + headers=validated_headers, + api_base=complete_api_base, + logging_obj=logging_obj, + client=client, + timeout=timeout, + ) + async def async_create_file( self, transformed_request: Union[bytes, str, dict, "TwoStepFileUploadConfig"], @@ -3742,6 +3857,20 @@ class BaseLLMHTTPHandler: if model is None: raise ValueError("model is required for create_batch") + if _is_async: + return self._avalidate_and_create_batch( + create_batch_data=create_batch_data, + litellm_params=litellm_params, + provider_config=provider_config, + headers=headers, + api_base=api_base, + api_key=api_key, + logging_obj=logging_obj, + client=client, + timeout=timeout, + model=model, + ) + headers = provider_config.validate_environment( api_key=api_key, headers=headers, @@ -3770,19 +3899,6 @@ class BaseLLMHTTPHandler: optional_params={}, ) - if _is_async: - return self.async_create_batch( - transformed_request=transformed_request, - litellm_params=litellm_params, - provider_config=provider_config, - headers=headers, - api_base=api_base, - logging_obj=logging_obj, - client=client, - timeout=timeout, - create_batch_data=create_batch_data, - ) - if client is None or not isinstance(client, HTTPHandler): sync_httpx_client = _get_httpx_client() else: @@ -3920,6 +4036,56 @@ class BaseLLMHTTPHandler: litellm_params=litellm_params, ) + async def _avalidate_and_create_batch( + self, + *, + create_batch_data: "CreateBatchRequest", + litellm_params: dict, # mutable-ok: mirrors the create_batch contract this dispatches for + provider_config: "BaseBatchesConfig", + headers: dict, # mutable-ok: mirrors the create_batch contract this dispatches for + api_base: str | None, + api_key: str | None, + logging_obj: "LiteLLMLoggingObj", + client: Union["HTTPHandler", "AsyncHTTPHandler"] | None, + timeout: float | httpx.Timeout | None, + model: str, + ) -> "LiteLLMBatch": + validated_headers: Final = await _avalidate_files_environment( + provider_config, + headers=headers, + model=model, + messages=[], + optional_params={}, + litellm_params=litellm_params, + api_key=api_key, + ) + complete_api_base: Final = provider_config.get_complete_batch_url( + api_base=api_base, + api_key=api_key, + model=model, + optional_params={}, + litellm_params=litellm_params, + data=create_batch_data, + ) + if not complete_api_base: + raise ValueError("api_base is required for create_batch") + return await self.async_create_batch( + transformed_request=provider_config.transform_create_batch_request( + model=model, + create_batch_data=create_batch_data, + litellm_params=litellm_params, + optional_params={}, + ), + litellm_params=litellm_params, + provider_config=provider_config, + headers=validated_headers, + api_base=complete_api_base, + logging_obj=logging_obj, + client=client, + timeout=timeout, + create_batch_data=create_batch_data, + ) + async def async_create_batch( self, transformed_request: bytes | str | dict, @@ -4521,7 +4687,8 @@ class BaseLLMHTTPHandler: ) # Validate environment and get headers - headers = provider_config.validate_environment( + headers = await _avalidate_files_environment( + provider_config, api_key=litellm_params.get("api_key"), headers=headers, model="", @@ -4646,7 +4813,8 @@ class BaseLLMHTTPHandler: ) # Validate environment and get headers - headers = provider_config.validate_environment( + headers = await _avalidate_files_environment( + provider_config, api_key=litellm_params.get("api_key"), headers=headers, model="", @@ -4770,7 +4938,8 @@ class BaseLLMHTTPHandler: ) # Validate environment and get headers - headers = provider_config.validate_environment( + headers = await _avalidate_files_environment( + provider_config, api_key=litellm_params.get("api_key"), headers=headers, model="", @@ -4958,7 +5127,7 @@ class BaseLLMHTTPHandler: else: async_httpx_client = client - prepared: Final = self._prepare_file_content_request( + prepared: Final = await self._aprepare_file_content_request( file_content_request=file_content_request, provider_config=provider_config, litellm_params=litellm_params, @@ -5004,7 +5173,7 @@ class BaseLLMHTTPHandler: client if client is not None else get_async_httpx_client(llm_provider=provider_config.custom_llm_provider) ) - prepared: Final = self._prepare_file_content_request( + prepared: Final = await self._aprepare_file_content_request( file_content_request=file_content_request, provider_config=provider_config, litellm_params=litellm_params, @@ -5062,16 +5231,31 @@ class BaseLLMHTTPHandler: optional_params={}, litellm_params=litellm_params, ) - logging_obj.pre_call( - input="", - api_key="", - additional_args={ - "api_base": url, - "headers": request_headers, - "file_id": file_content_request.get("file_id"), - }, + return _logged_file_content_request(url, params, request_headers, file_content_request, logging_obj) + + @staticmethod + async def _aprepare_file_content_request( + file_content_request: "FileContentRequest", + provider_config: BaseFilesConfig, + litellm_params: dict, + headers: dict, + logging_obj: LiteLLMLoggingObj, + ) -> "_PreparedFileContentRequest": + url, params = provider_config.transform_file_content_request( + file_content_request=file_content_request, + optional_params={}, + litellm_params=litellm_params, ) - return _PreparedFileContentRequest(url=url, params=params, headers=request_headers) + request_headers: Final = await _avalidate_files_environment( + provider_config, + api_key=litellm_params.get("api_key"), + headers=headers, + model="", + messages=[], + optional_params={}, + litellm_params=litellm_params, + ) + return _logged_file_content_request(url, params, request_headers, file_content_request, logging_obj) def _prepare_fake_stream_request( self, diff --git a/litellm/llms/openai/chat/gpt_transformation.py b/litellm/llms/openai/chat/gpt_transformation.py index 6204bed8109..e7a7cda0b67 100644 --- a/litellm/llms/openai/chat/gpt_transformation.py +++ b/litellm/llms/openai/chat/gpt_transformation.py @@ -58,6 +58,7 @@ from litellm.types.utils import ( from litellm.utils import convert_to_model_response_object from ..common_utils import OpenAIError +from ..workload_identity import get_workload_identity_bearer_token, resolve_openai_workload_identity_config if TYPE_CHECKING: from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj @@ -73,6 +74,11 @@ else: _NO_TOOLS_UPDATE: Final[Mapping[str, object]] = MappingProxyType({}) +def _litellm_params_str(litellm_params: Mapping[str, object] | None, key: str) -> str | None: + value: Final = litellm_params.get(key) if litellm_params is not None else None + return value if isinstance(value, str) else None + + class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): """ Reference: https://platform.openai.com/docs/api-reference/chat/create @@ -765,28 +771,39 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig): """ Calls OpenAI's `/v1/models` endpoint and returns the list of models. """ - - if api_base is None: - api_base = "https://api.openai.com" - if api_key is None: - api_key = get_secret_str("OPENAI_API_KEY") - - # Strip api_base to just the base URL (scheme + host + port) - parsed_url: Final = httpx.URL(api_base) - base_url = f"{parsed_url.scheme}://{parsed_url.host}" - if parsed_url.port: - base_url += f":{parsed_url.port}" - - response: Final = litellm.module_level_client.get( - url=f"{base_url}/v1/models", - headers={"Authorization": f"Bearer {api_key}"}, + return self._fetch_model_ids( + api_base=api_base, bearer_token=get_secret_str("OPENAI_API_KEY") if api_key is None else api_key ) + def discover_models( + self, litellm_params: Mapping[str, object] | None = None + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + if type(self) is not OpenAIGPTConfig: + return super().discover_models(litellm_params) + api_key: Final = _litellm_params_str(litellm_params, "api_key") + api_base: Final = _litellm_params_str(litellm_params, "api_base") + workload_identity_config: Final = resolve_openai_workload_identity_config( + api_key=api_key, api_base=api_base, litellm_params=litellm_params + ) + if workload_identity_config is None: + return self.get_models(api_key=api_key, api_base=api_base) + return self._fetch_model_ids( + api_base=api_base, bearer_token=get_workload_identity_bearer_token(workload_identity_config) + ) + + @staticmethod + def _fetch_model_ids( + api_base: str | None, bearer_token: str | None + ) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override + parsed_url: Final = httpx.URL(api_base or "https://api.openai.com") + port_suffix: Final = f":{parsed_url.port}" if parsed_url.port else "" + response: Final = litellm.module_level_client.get( + url=f"{parsed_url.scheme}://{parsed_url.host}{port_suffix}/v1/models", + headers={"Authorization": f"Bearer {bearer_token}"}, + ) if response.status_code != 200: raise Exception(f"Failed to get models: {response.text}") - - models: Final = response.json()["data"] - return [model["id"] for model in models] + return [model["id"] for model in response.json()["data"]] @staticmethod def get_api_key(api_key: str | None = None) -> str | None: diff --git a/litellm/llms/openai/openai.py b/litellm/llms/openai/openai.py index e3792fe9dfa..2461fb4e1af 100644 --- a/litellm/llms/openai/openai.py +++ b/litellm/llms/openai/openai.py @@ -382,8 +382,11 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization: str | None = None, client: OpenAI | AsyncOpenAI | None = None, shared_session: Optional["ClientSession"] = None, + litellm_params: Mapping[str, object] | None = None, ) -> OpenAI | AsyncOpenAI | None: - workload_identity_config: Final = resolve_openai_workload_identity_config(api_key=api_key, api_base=api_base) + workload_identity_config: Final = resolve_openai_workload_identity_config( + api_key=api_key, api_base=api_base, litellm_params=litellm_params + ) client_initialization_params: Final[dict] = locals() if client is None: if not isinstance(max_retries, int): @@ -773,6 +776,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, organization=organization, stream_options=stream_options, + litellm_params=litellm_params, ) else: if not isinstance(max_retries, int): @@ -786,6 +790,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, organization=organization, client=client, + litellm_params=litellm_params, ) ## LOGGING @@ -928,6 +933,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization=organization, client=client, shared_session=shared_session, + litellm_params=litellm_params, ) ## LOGGING @@ -1024,6 +1030,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=None, headers=None, stream_options: dict | None = None, + litellm_params: Mapping[str, object] | None = None, ): data["stream"] = True data.update(self.get_stream_options(stream_options=stream_options, api_base=api_base)) @@ -1037,6 +1044,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, organization=organization, client=client, + litellm_params=litellm_params, ) ## LOGGING logging_obj.pre_call( @@ -1109,6 +1117,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): organization=organization, client=client, shared_session=shared_session, + litellm_params=litellm_params, ) ## LOGGING logging_obj.pre_call( @@ -1243,6 +1252,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): client: AsyncOpenAI | None = None, max_retries=None, shared_session: Optional["ClientSession"] = None, + litellm_params: Mapping[str, object] | None = None, ): try: openai_aclient: Final[AsyncOpenAI] = self._get_openai_client( @@ -1253,6 +1263,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): max_retries=max_retries, client=client, shared_session=shared_session, + litellm_params=litellm_params, ) raw_response: Final = await self.make_openai_embedding_request( openai_aclient=openai_aclient, @@ -1316,6 +1327,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): aembedding=None, max_retries: int | None = None, shared_session: Optional["ClientSession"] = None, + litellm_params: Mapping[str, object] | None = None, ) -> EmbeddingResponse: super().embedding() try: @@ -1342,6 +1354,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): client=client, max_retries=max_retries, shared_session=shared_session, + litellm_params=litellm_params, ) openai_client: Final[OpenAI] = self._get_openai_client( @@ -1351,6 +1364,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM): timeout=timeout, max_retries=max_retries, client=client, + litellm_params=litellm_params, ) ## embedding CALL diff --git a/litellm/llms/openai/responses/transformation.py b/litellm/llms/openai/responses/transformation.py index 67b9e157832..1452ecebca9 100644 --- a/litellm/llms/openai/responses/transformation.py +++ b/litellm/llms/openai/responses/transformation.py @@ -30,6 +30,7 @@ from litellm.types.llms.openai import * from litellm.types.responses.main import * from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import LlmProviders +from litellm.types.workload_identity import OPENAI_WIF_KWARGS_KEYS from ..common_utils import OpenAIError from ..workload_identity import get_workload_identity_bearer_token, resolve_openai_workload_identity_config @@ -599,7 +600,11 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig): api_key = litellm_params.api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") headers.setdefault("Content-Type", "application/json") workload_identity_config: Final = ( - resolve_openai_workload_identity_config(api_key=api_key, api_base=litellm_params.api_base) + resolve_openai_workload_identity_config( + api_key=api_key, + api_base=litellm_params.api_base, + litellm_params=litellm_params.model_dump(include=set(OPENAI_WIF_KWARGS_KEYS)), + ) if self.custom_llm_provider is LlmProviders.OPENAI else None ) diff --git a/litellm/llms/openai/workload_identity.py b/litellm/llms/openai/workload_identity.py index 369b4f1e3f7..e2971a2c100 100644 --- a/litellm/llms/openai/workload_identity.py +++ b/litellm/llms/openai/workload_identity.py @@ -1,5 +1,6 @@ from __future__ import annotations +from collections.abc import Mapping from dataclasses import dataclass from functools import lru_cache from typing import TYPE_CHECKING, Final @@ -43,6 +44,7 @@ class OpenAIWorkloadIdentityConfig: def resolve_openai_workload_identity_config( api_key: str | None, api_base: str | None, + litellm_params: Mapping[str, object] | None = None, ) -> OpenAIWorkloadIdentityConfig | None: static_api_key: Final = normalize_nonempty_secret_str(api_key) or normalize_nonempty_secret_str( get_secret_str("OPENAI_API_KEY") @@ -54,10 +56,12 @@ def resolve_openai_workload_identity_config( ) if not _targets_openai_api(effective_api_base): return None - identity_provider_id: Final = get_secret_str("OPENAI_IDENTITY_PROVIDER_ID") - service_account_id: Final = get_secret_str("OPENAI_SERVICE_ACCOUNT_ID") - token_file: Final = get_secret_str("OPENAI_IDENTITY_TOKEN_FILE") - if not identity_provider_id or not service_account_id or not token_file: + identity_provider_id: Final = _config_value( + litellm_params, "openai_identity_provider_id", "OPENAI_IDENTITY_PROVIDER_ID" + ) + service_account_id: Final = _config_value(litellm_params, "openai_service_account_id", "OPENAI_SERVICE_ACCOUNT_ID") + token_file: Final = _config_value(litellm_params, "openai_identity_token_file", "OPENAI_IDENTITY_TOKEN_FILE") + if identity_provider_id is None or service_account_id is None or token_file is None: return None return OpenAIWorkloadIdentityConfig( identity_provider_id=identity_provider_id, @@ -77,6 +81,13 @@ async def get_workload_identity_bearer_token_for_api_base(api_base: str) -> str return await _workload_identity_auth(config).get_token_async() +def _config_value(litellm_params: Mapping[str, object] | None, param_key: str, env_name: str) -> str | None: + param_value: Final = litellm_params.get(param_key) if litellm_params is not None else None + if isinstance(param_value, str) and param_value: + return param_value + return normalize_nonempty_secret_str(get_secret_str(env_name)) + + def _targets_openai_api(api_base: str | None) -> bool: if api_base is None: return True diff --git a/litellm/main.py b/litellm/main.py index 6cfc2b8af55..122de9a02c0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -144,6 +144,7 @@ from litellm.types.utils import ( RawRequestTypedDict, StreamingChoices, ) +from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS, OPENAI_WIF_KWARGS_KEYS from litellm.utils import ( Choices, CustomStreamWrapper, @@ -5489,7 +5490,9 @@ def completion( api_base=api_base, api_key=api_key, litellm_params=( - GenericLiteLLMParams(**_supplemental_provider_params) if _supplemental_provider_params else None + GenericLiteLLMParams.model_validate(_supplemental_provider_params) + if _supplemental_provider_params + else None ), ) @@ -5716,7 +5719,12 @@ def completion( gigachat_access_token=kwargs.get("gigachat_access_token"), **{ key: kwargs[key] - for key in (*AWS_CREDENTIAL_KWARGS_KEYS, PROVIDER_AFFINITY_HEADER_KWARG_KEY) + for key in ( + *AWS_CREDENTIAL_KWARGS_KEYS, + *ANTHROPIC_WIF_KWARGS_KEYS, + *OPENAI_WIF_KWARGS_KEYS, + PROVIDER_AFFINITY_HEADER_KWARG_KEY, + ) if key in kwargs }, ) @@ -6553,6 +6561,7 @@ def embedding( aembedding=aembedding, max_retries=max_retries, shared_session=shared_session, + litellm_params=litellm_params_dict, ) elif custom_llm_provider == "databricks": api_base = api_base or litellm.api_base or get_secret("DATABRICKS_API_BASE") @@ -7802,7 +7811,7 @@ async def amoderation( # only supports open ai for now api_key = api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY") - optional_params: Final = GenericLiteLLMParams(**kwargs) + optional_params: Final = GenericLiteLLMParams.model_validate(kwargs) litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None) _dynamic_api_base = None try: @@ -8519,7 +8528,7 @@ def speech( VertexAITextToSpeechConfig, ) - generic_optional_params: Final = GenericLiteLLMParams(**kwargs) + generic_optional_params: Final = GenericLiteLLMParams.model_validate(kwargs) # Handle Gemini models separately (they use speech_to_completion_bridge) if "gemini" in model: diff --git a/litellm/models/credentials.py b/litellm/models/credentials.py index 0878eea5769..e96e406d98e 100644 --- a/litellm/models/credentials.py +++ b/litellm/models/credentials.py @@ -7,7 +7,7 @@ layer; ``litellm.types.utils`` re-exports them for backwards compatibility. from collections.abc import Mapping -from pydantic import BaseModel, model_validator +from pydantic import BaseModel, Field, model_validator class CredentialBase(BaseModel): @@ -17,6 +17,10 @@ class CredentialBase(BaseModel): class CredentialItem(CredentialBase): credential_values: dict + # PATCH-only instruction naming keys to drop from the stored credential_values. It describes an + # edit rather than the credential, so it stays out of dumps: those feed config loading, the DB + # write, and the in-memory list, none of which have a place for it. + credential_values_to_delete: tuple[str, ...] | None = Field(default=None, exclude=True) class CreateCredentialItem(CredentialBase): @@ -36,3 +40,4 @@ class UpdateCredentialItem(BaseModel): credential_info: Mapping[str, object] credential_values: Mapping[str, object] | None = None model_id: str | None = None + credential_values_to_delete: tuple[str, ...] | None = None diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 7ef82009184..c3c1032a3f0 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -20,6 +20,7 @@ from litellm.constants import ( MINIMUM_CUSTOM_KEY_LENGTH, STANDARD_CUSTOMER_ID_HEADERS, ) +from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.litellm_core_utils.url_utils import ( SSRFError, @@ -34,7 +35,8 @@ from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_me from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_ENDPOINT_MARKER, ) -from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, Deployment +from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, Deployment, server_owned_wif_fields_present +from litellm.types.router import reject_server_owned_wif_params as _reject_server_owned_wif_params from litellm.types.utils import CustomPricingLiteLLMParams @@ -228,6 +230,63 @@ def _allow_model_level_clientside_configurable_parameters( # ``extra_body.aws_web_identity_token``) without re-validating, so the # banned-key check has to descend into it the same way it descends into # ``litellm_embedding_config``. +# Re-exported from litellm.types.router, where it lives so the router can call it on a +# post-authentication merge without core importing from the proxy package. +reject_server_owned_wif_params = _reject_server_owned_wif_params + +_CREDENTIAL_VALUES: Final = TypeAdapter(dict[str, object]) + + +def reject_federated_credential_reference(body: Mapping[str, object]) -> None: + """Raise ``ValueError`` if a request body picks the federated identity by credential name. + + ``load_credentials_from_list`` merges a named credential's values into the call, so a body + naming a federated credential moves the token exchange onto that credential's federation rule + and organization exactly as sending the federation fields inline would, which + ``reject_server_owned_wif_params`` already refuses. Attaching one is a deployment decision, so + it is refused with no client-side opt-in to relax it, the same as the inline form. + + Only credentials already loaded into memory can be resolved here, which is every credential the + proxy would resolve for the call itself: ``load_credentials_from_list`` reads the same list. + ``route_manages_deployments`` names the routes this is not applied to, where choosing the + credential a deployment federates through is the point of the call. + """ + named: Final = body.get("litellm_credential_name") + if not isinstance(named, str) or not named: + return + wif_fields: Final = server_owned_wif_fields_present( + _CREDENTIAL_VALUES.validate_python(CredentialAccessor.get_credential_values(named)) + ) + if wif_fields: + raise ValueError( + f"Rejected Request: litellm_credential_name={named!r} names a credential configured for " + f"workload identity federation ({wif_fields[0]}), which a request body cannot choose. " + "A proxy admin attaches it to a deployment." + ) + + +_DEPLOYMENT_MANAGEMENT_ROUTES: Final[frozenset[str]] = frozenset( + ("/model/new", "/model/update", "/model/delete", "/health/test_connection") +) +_DEPLOYMENT_ID_UPDATE_ROUTE: Final = re.compile(r"^/model/[^/]+/update$") + + +def route_manages_deployments(route: str | None) -> bool: + """Whether ``route`` configures a deployment instead of calling one. + + These are the routes that reach ``ModelManagementAuthChecks.can_user_make_model_call``, where + a federated write is judged by ``_reject_non_admin_wif_write`` against what the write sets and + what the deployment already stores: a proxy admin goes through, anyone else is refused with a + 403 naming the field. ``reject_federated_credential_reference`` runs ahead of that gate on + every route, so without this exemption a proxy admin could not attach a federated credential + to a deployment over the API or the Admin UI at all, leaving a static ``config.yaml`` entry as + the only way to configure the feature the rejection tells the caller to go configure. + """ + return route is not None and ( + route in _DEPLOYMENT_MANAGEMENT_ROUTES or _DEPLOYMENT_ID_UPDATE_ROUTE.fullmatch(route) is not None + ) + + _NESTED_CONFIG_KEYS: Final[tuple[str, ...]] = ("litellm_embedding_config", "extra_body") # Metadata containers that carry per-request configuration consumed by the @@ -380,12 +439,17 @@ def _check_banned_params( general_settings: dict, llm_router: Router | None, model: str, + *, + manages_deployments: bool = False, ) -> None: """Raise ``ValueError`` if ``body`` carries a banned param without admin opt-in. Shared between the root-level check and the nested-config check so a new banned param only needs to be added in one place. """ + reject_server_owned_wif_params(body) + if not manages_deployments: + reject_federated_credential_reference(body) for param in _BANNED_REQUEST_BODY_PARAMS: if param not in body: continue @@ -472,7 +536,14 @@ def _reject_url_valued_fallback_target(value: str) -> None: ) -def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: Router | None, model: str) -> bool: +def is_request_body_safe( + request_body: dict, + general_settings: dict, + llm_router: Router | None, + model: str, + *, + route: str | None = None, +) -> bool: """ Check if the request body is safe. @@ -500,25 +571,27 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: """ if "model_list" in request_body: raise ValueError("Rejected Request: model_list is not allowed in the request body.") - _check_banned_params(request_body, general_settings, llm_router, model) + manages_deployments: Final = route_manages_deployments(route) + _check_banned_params(request_body, general_settings, llm_router, model, manages_deployments=manages_deployments) for nested_key in _NESTED_CONFIG_KEYS: nested = _coerce_metadata_to_dict(request_body.get(nested_key)) if nested is not None: - _check_banned_params(nested, general_settings, llm_router, model) + _check_banned_params(nested, general_settings, llm_router, model, manages_deployments=manages_deployments) for metadata_key in _NESTED_METADATA_KEYS: metadata = _coerce_metadata_to_dict(request_body.get(metadata_key)) if metadata is not None: - _check_banned_params(metadata, general_settings, llm_router, model) + _check_banned_params(metadata, general_settings, llm_router, model, manages_deployments=manages_deployments) if any(isinstance(key, str) and key.startswith(f"{metadata_key}[") for key in request_body): _check_banned_params( extract_nested_form_metadata(form_data=request_body, prefix=f"{metadata_key}["), general_settings, llm_router, model, + manages_deployments=manages_deployments, ) for target in iter_request_fallback_targets(request_body): if isinstance(target, dict): - _check_banned_params(target, general_settings, llm_router, model) + _check_banned_params(target, general_settings, llm_router, model, manages_deployments=manages_deployments) target_model = target.get("model") if isinstance(target_model, str): _reject_url_valued_fallback_target(target_model) @@ -526,6 +599,9 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: _reject_url_valued_fallback_target(target) litellm_params: Final = _coerce_metadata_to_dict(request_body.get("litellm_params")) if litellm_params is not None: + reject_server_owned_wif_params(litellm_params) + if not manages_deployments: + reject_federated_credential_reference(litellm_params) litellm_params_metadata: Final = _coerce_metadata_to_dict(litellm_params.get("metadata")) if litellm_params_metadata is not None: _check_banned_params( @@ -533,6 +609,7 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: general_settings, llm_router, model, + manages_deployments=manages_deployments, ) return True @@ -585,6 +662,7 @@ async def pre_db_read_auth_checks( general_settings=general_settings, llm_router=llm_router, model=request_data.get("model", ""), # [TODO] use model passed in url as well (azure openai routes) + route=route, ) # Check 3. Check if IP address is allowed diff --git a/litellm/proxy/common_utils/credential_hydration.py b/litellm/proxy/common_utils/credential_hydration.py new file mode 100644 index 00000000000..2aabe8cad5c --- /dev/null +++ b/litellm/proxy/common_utils/credential_hydration.py @@ -0,0 +1,189 @@ +"""Shared helper for resolving a named Credential's values server-side. + +Memory first (``litellm.credential_list``, already decrypted -- matching +``CredentialAccessor.get_credential_values``), then a DB decrypt fallback for a pod whose +in-memory list has not yet picked up a credential another pod just wrote or updated. +""" + +import asyncio +from collections.abc import Mapping +from itertools import chain +from types import MappingProxyType +from typing import Final + +import litellm +from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper +from litellm.proxy.utils import PrismaClient +from litellm.repositories.credentials_repository import CredentialsRepository +from litellm.router_utils.clientside_credential_handler import clientside_credential_keys +from litellm.types.router import ( + GenericLiteLLMParams, + server_owned_wif_fields_named, + server_owned_wif_fields_present, +) +from litellm.types.utils import CredentialItem, LlmProviders, server_owned_wif_litellm_params + +_LITELLM_PROVIDER_IDS: Final = frozenset(provider.value for provider in LlmProviders) + +_FEDERATION_SURFACE_FIELDS: Final = frozenset( + ( + *clientside_credential_keys, + "configurable_clientside_auth_params", + "litellm_credential_name", + *server_owned_wif_litellm_params, + ) +) + + +def write_touches_federation_surface(incoming: Mapping[str, object] | None) -> bool: + """Whether this write can move or re-scope the token a federated deployment mints. + + Three groups of fields can. The federation parameters choose which server-side secret is read + and what the minted token is scoped to. ``litellm_credential_name`` resolves to those same + parameters by reference. ``api_key``, ``api_base``, ``base_url``, and the + ``configurable_clientside_auth_params`` that let a caller override them decide where the + resulting token is sent. A write setting none of them leaves the federation configuration + exactly as the proxy admin left it, so renaming a federated deployment or changing its rpm + stays an ordinary team-admin edit. + """ + return incoming is not None and not _FEDERATION_SURFACE_FIELDS.isdisjoint(incoming.keys()) + + +def stored_credential_provider(credential_provider: object) -> str | None: + """The dashboard stores its display casing (``Anthropic``) on credentials it creates, so the + provider a credential names is the lowercased value when that is a litellm provider id.""" + if not isinstance(credential_provider, str): + return None + lowered: Final = credential_provider.lower() + return lowered if lowered in _LITELLM_PROVIDER_IDS else None + + +def decrypted_or_stored(key: str, value: str) -> str: + """The stored value decrypted, or as stored when it was never encrypted (a config.yaml value).""" + decrypted: Final = decrypt_value_helper(value=value, key=key) + return value if decrypted is None else decrypted + + +def _decrypted(db_credential: CredentialItem) -> CredentialItem: + """The stored credential with every value decrypted, leaving already-plaintext values alone.""" + decrypted_values: Final = MappingProxyType( + {key: decrypted_or_stored(key, value) for key, value in db_credential.credential_values.items()} + ) + return CredentialItem( + credential_name=db_credential.credential_name, + credential_values=decrypted_values, # pyright: ignore[reportArgumentType] # declared dict[str, str], and pydantic copies this mapping into one on validation; LIT002 rules out building that dict here + credential_info=db_credential.credential_info, + ) + + +async def hydrate_named_credential_authoritative( + credential_name: str, + prisma_client: PrismaClient | None, +) -> CredentialItem | None: + """The stored credential, preferring the row over this pod's in-memory copy. + + ``hydrate_named_credential`` reads memory first, which is right when serving a request. A + management operation cannot: on a pod whose in-memory copy predates another pod's update, it + would export the superseded JWKS, or discover models against superseded values. Same reason + ``named_credential_wif_fields`` reads both. + """ + if prisma_client is None: + return await hydrate_named_credential(credential_name, prisma_client) + db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name) + if db_credential is None: + return await hydrate_named_credential(credential_name, prisma_client) + return _decrypted(db_credential) + + +async def hydrate_named_credential( + credential_name: str, + prisma_client: PrismaClient | None, +) -> CredentialItem | None: + for credential in litellm.credential_list: + if credential.credential_name == credential_name: + return credential + if prisma_client is None: + return None + db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name) + if db_credential is None: + return None + return _decrypted(db_credential) + + +async def named_credential_wif_fields( + credential_name: str, + prisma_client: PrismaClient | None, +) -> tuple[str, ...]: + """Federation field names a write to ``credential_name`` would touch, from memory AND the row. + + Resolution reads memory first and stops there, which is right when serving a request. An + authorization decision cannot: a pod whose in-memory copy predates an admin adding federation + fields would see none and allow the write. This reads both and returns the union, so the gate + refuses whenever either side says the credential is server-owned. + """ + matching: Final = tuple(c for c in litellm.credential_list if c.credential_name == credential_name) + in_memory: Final = tuple(chain.from_iterable(server_owned_wif_fields_named(c.credential_values) for c in matching)) + if prisma_client is None: + return in_memory + db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name) + stored: Final = () if db_credential is None else server_owned_wif_fields_named(db_credential.credential_values) + return tuple(dict.fromkeys(in_memory + stored)) + + +def submitted_litellm_params(params: GenericLiteLLMParams | None) -> Mapping[str, object] | None: + """The fields a pydantic write actually set, as the mapping the federation gate reads. + + Only the set fields belong here: ``GenericLiteLLMParams`` declares every federation field, so + the whole model would report every write as touching all of them. + """ + if params is None: + return None + return MappingProxyType({name: getattr(params, name, None) for name in params.model_fields_set}) + + +async def effective_server_owned_wif_fields( + stored: Mapping[str, object] | None, + incoming: Mapping[str, object] | None, + prisma_client: PrismaClient | None, +) -> tuple[str, ...]: + """Federation field names the deployment would carry AFTER this write. + + Authorization has to read the resulting deployment, not the submitted payload. A patch that + names no federation field still lands on a deployment that has them, and a patch that only + attaches ``litellm_credential_name`` inherits whatever that credential holds. + + The two sides are matched differently on purpose. ``stored`` is matched by VALUE, since it is + a full deployment and a declared-but-unset field is not a federation field it carries. + ``incoming`` holds only the keys the write actually set, so an explicit null still counts as + touching the field. + """ + from_stored: Final = () if stored is None else server_owned_wif_fields_present(stored) + from_incoming: Final = () if incoming is None else server_owned_wif_fields_named(incoming.keys()) + from_credential: Final = tuple( + chain.from_iterable( + await asyncio.gather( + *( + named_credential_wif_fields(credential_name, prisma_client) + for credential_name in _effective_credential_names(stored, incoming) + ) + ) + ) + ) + return tuple(dict.fromkeys(from_stored + from_incoming + from_credential)) + + +def _effective_credential_names( + stored: Mapping[str, object] | None, + incoming: Mapping[str, object] | None, +) -> tuple[str, ...]: + """Both the credential the deployment already carries and the one this write names. + + Taking only the incoming name would let a write clear its way out: detaching a federated + credential, by sending ``litellm_credential_name: null`` alongside an api_key or api_base of + the caller's choosing, would leave nothing federated to find and the write would be allowed. + Detaching an administrator's federated credential is itself an administrator's action, so the + stored name counts whatever the write says. + """ + from_stored: Final = None if stored is None else stored.get("litellm_credential_name") + from_incoming: Final = None if incoming is None else incoming.get("litellm_credential_name") + return tuple(dict.fromkeys(name for name in (from_stored, from_incoming) if isinstance(name, str))) diff --git a/litellm/proxy/credential_endpoints/endpoints.py b/litellm/proxy/credential_endpoints/endpoints.py index f99cce14722..af2b7f0df8a 100644 --- a/litellm/proxy/credential_endpoints/endpoints.py +++ b/litellm/proxy/credential_endpoints/endpoints.py @@ -9,26 +9,134 @@ from typing import ( cast, # noqa: TID251 # jsonify_object in proxy/utils.py is annotated with a bare dict ) -from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response +from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response, status from pydantic import TypeAdapter import litellm from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.litellm_logging import _get_masked_values +from litellm.llms.anthropic.wif import ( + ExportedJwks, + NotAnInternalIssuerCredential, + UnbuildableIdentitySource, + anthropic_internal_issuer_jwks, +) from litellm.models.credentials import UpdateCredentialItem -from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy._types import ( + CommonProxyErrors, + LitellmUserRoles, + ProxyErrorTypes, + ProxyException, + UserAPIKeyAuth, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.credential_hydration import ( + hydrate_named_credential, + hydrate_named_credential_authoritative, + named_credential_wif_fields, + stored_credential_provider, +) from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper from litellm.proxy.utils import handle_exception_on_proxy, jsonify_object from litellm.repositories.base_repository import is_unique_violation from litellm.repositories.credentials_repository import CredentialsRepository +from litellm.types.router import server_owned_wif_fields_named from litellm.types.utils import CreateCredentialItem, CredentialItem router: Final = APIRouter() _CREDENTIAL_DICT_ADAPTER: Final = TypeAdapter(dict[str, object]) +def _reject_non_admin_wif_fields( + wif_fields: tuple[str, ...], + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """A credential referenced by ``litellm_credential_name`` feeds its values into the same + workload identity federation resolution as a deployment's own ``litellm_params``. Only proxy + admins may touch a server-owned WIF field, whether they write it, drop it, or edit a stored + credential that already carries one. + """ + if not wif_fields or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: + return + raise ProxyException( + message=( + f"Only proxy admins can change {wif_fields[0]!r}, a server-owned workload identity federation parameter." + ), + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param=wif_fields[0], + ) + + +def _incoming_wif_fields(incoming_values: Mapping[str, object], credential: UpdateCredentialItem) -> tuple[str, ...]: + """WIF fields the request touches: the ones its values set (to any value, ``None`` included, + since the key alone is what the federation resolver reacts to), whether the caller sent them + or named a deployment through ``model_id`` for the proxy to copy them from, plus the ones it + names in ``credential_values_to_delete``, since dropping a federation field off the stored + credential breaks every deployment referencing it just as installing one would redirect them. + """ + return server_owned_wif_fields_named(incoming_values) + server_owned_wif_fields_named( + credential.credential_values_to_delete or () + ) + + +def _stored_wif_fields(stored_credential: CredentialItem) -> tuple[str, ...]: + return server_owned_wif_fields_named(stored_credential.credential_values) + + +def _reject_overlapping_credential_values(credential: UpdateCredentialItem) -> None: + overlap: Final = frozenset(credential.credential_values or ()) & frozenset( + credential.credential_values_to_delete or () + ) + if overlap: + raise HTTPException( + status_code=400, + detail=f"credential_values_to_delete overlaps credential_values for key(s): {sorted(overlap)}", + ) + + +def _without_null_values(credential_values: Mapping[str, object]) -> dict[str, object]: + """A null carries no credential, and the federation resolver refuses a foreign variant's field by + KEY, so a stored ``{"anthropic_issuer_url": null}`` wedges every deployment that names this + credential. ``model_dump(exclude_none=True)`` cannot do this: it drops the model's own null + fields, and ``credential_values`` is a mapping inside one of them. + """ + return {key: value for key, value in credential_values.items() if value is not None} + + +def _sync_in_memory_credential(credential: CredentialItem, credential_name: str, new_name: str) -> None: + """Mirror a DB credential update into the in-memory ``credential_list`` used by request-time + resolution; a no-op if the credential isn't loaded in memory (e.g. proxy restarted since boot). + """ + existing_in_memory: CredentialItem | None = None + for cred in litellm.credential_list: + if cred.credential_name == credential_name: + existing_in_memory = cred + break + + if existing_in_memory is None: + return + + in_memory_values: Final = dict(existing_in_memory.credential_values or {}) + if credential.credential_values: + in_memory_values.update(_without_null_values(credential.credential_values)) + for key in credential.credential_values_to_delete or (): + in_memory_values.pop(key, None) + in_memory_info: Final = dict(existing_in_memory.credential_info or {}) + if credential.credential_info: + in_memory_info.update(credential.credential_info) + updated_in_memory: Final = CredentialItem( + credential_name=new_name, + credential_values=in_memory_values, + credential_info=in_memory_info, + ) + # Remove old entry if renamed, then use upsert_credentials to handle duplicates + if new_name != credential_name: + litellm.credential_list = [c for c in litellm.credential_list if c.credential_name != credential_name] + CredentialAccessor.upsert_credentials([updated_in_memory]) + + class CredentialHelperUtils: @staticmethod def encrypt_credential_values(credential: CredentialItem, new_encryption_key: str | None = None) -> CredentialItem: @@ -108,13 +216,17 @@ async def create_credential( status_code=400, detail="Credential values are required. Unable to infer credential values from model ID.", ) + _reject_non_admin_wif_fields(server_owned_wif_fields_named(credential_values), user_api_key_dict) + _reject_non_admin_wif_fields( + await named_credential_wif_fields(credential.credential_name, prisma_client), user_api_key_dict + ) processed_credential: Final = CredentialItem( credential_name=credential.credential_name, - credential_values=_CREDENTIAL_DICT_ADAPTER.validate_python(credential_values), + credential_values=_without_null_values(_CREDENTIAL_DICT_ADAPTER.validate_python(credential_values)), credential_info=credential.credential_info, ) encrypted_credential: Final = CredentialHelperUtils.encrypt_credential_values(processed_credential) - credentials_dict: Final = encrypted_credential.model_dump() + credentials_dict: Final = encrypted_credential.model_dump(exclude_none=True) credentials_dict_jsonified: Final = cast( # cast-ok: deep-copies a model_dump, so keys are str "dict[str, object]", jsonify_object(credentials_dict) ) @@ -204,6 +316,65 @@ async def get_credential_by_name( raise handle_exception_on_proxy(e) +@router.get( + "/credentials/{credential_name:path}/jwks", + dependencies=(Depends(user_api_key_auth),), + tags=["credential management"], +) +async def get_credential_internal_issuer_jwks( + credential_name: str = Path(..., description="The credential name, percent-decoded; may contain slashes"), + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI resolves the dependency from the default +): + """ + Export the public JWKS for an anthropic ``internal_issuer`` credential, so the operator can + register it on the Anthropic federation issuer from the UI. Never touches the private signing + key: only its derived public JWKS leaves this process. 404s for any other credential shape. + """ + from litellm.proxy.proxy_server import prisma_client + + if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise HTTPException( + status_code=403, + detail={"error": "Only proxy admins can export a credential's JWKS."}, + ) + + try: + credential: Final = await hydrate_named_credential_authoritative(credential_name, prisma_client) + credential_provider: Final = ( + None + if credential is None + else stored_credential_provider(credential.credential_info.get("custom_llm_provider")) + ) + if credential is None or credential_provider != "anthropic": + raise HTTPException( + status_code=404, + detail={"error": f"No anthropic credential named {credential_name!r}."}, + ) + match anthropic_internal_issuer_jwks(credential.credential_values): + case ExportedJwks(document): + return Response(content=document, media_type="application/json") + case NotAnInternalIssuerCredential(required_param, required_value): + raise HTTPException( + status_code=404, + detail={ + "error": ( + f"Credential {credential_name!r} is not configured with " + f"{required_param}={required_value!r}." + ) + }, + ) + case UnbuildableIdentitySource(message): + raise HTTPException( + status_code=400, + detail={"error": message}, + ) + except HTTPException: + raise + except Exception as e: # noqa: BLE001 # endpoint boundary: every failure becomes the proxy's error contract + verbose_proxy_logger.exception(e) + raise handle_exception_on_proxy(e) + + @router.get( "/credentials/by_model/{model_id}", dependencies=[Depends(user_api_key_auth)], @@ -268,6 +439,9 @@ async def delete_credential( status_code=500, detail={"error": CommonProxyErrors.db_not_connected_error.value}, ) + _reject_non_admin_wif_fields( + await named_credential_wif_fields(credential_name, prisma_client), user_api_key_dict + ) deleted: Final = await CredentialsRepository(prisma_client).delete_by_name(credential_name) if deleted is None: raise HTTPException( @@ -307,9 +481,10 @@ def update_db_credential( # update litellm params if encrypted_credential.credential_values: # Encrypt any sensitive values - encrypted_params: Final = {k: v for k, v in encrypted_credential.credential_values.items()} + merged_credential.credential_values.update(_without_null_values(encrypted_credential.credential_values)) - merged_credential.credential_values.update(encrypted_params) + for key in updated_patch.credential_values_to_delete or (): + merged_credential.credential_values.pop(key, None) # update model info if encrypted_credential.credential_info: @@ -340,6 +515,13 @@ async def update_credential( from litellm.proxy.proxy_server import prisma_client try: + _reject_overlapping_credential_values(credential) + incoming_values: Final = _CREDENTIAL_DICT_ADAPTER.validate_python( + _resolve_deployment_credentials(llm_router, credential.model_id) + if credential.model_id + else credential.credential_values or {} + ) + _reject_non_admin_wif_fields(_incoming_wif_fields(incoming_values, credential), user_api_key_dict) if prisma_client is None: raise HTTPException( status_code=500, @@ -349,18 +531,20 @@ async def update_credential( db_credential: Final = await credentials_repository.find_by_name(credential_name) if db_credential is None: raise HTTPException(status_code=404, detail="Credential not found in DB.") + _reject_non_admin_wif_fields(_stored_wif_fields(db_credential), user_api_key_dict) + if credential.credential_name != credential_name: + shadowed_credential: Final = await hydrate_named_credential(credential.credential_name, prisma_client) + if shadowed_credential is not None: + _reject_non_admin_wif_fields(_stored_wif_fields(shadowed_credential), user_api_key_dict) patch: Final = CredentialItem( credential_name=credential.credential_name, credential_info=_CREDENTIAL_DICT_ADAPTER.validate_python(credential.credential_info), - credential_values=_CREDENTIAL_DICT_ADAPTER.validate_python( - _resolve_deployment_credentials(llm_router, credential.model_id) - if credential.model_id - else credential.credential_values or {} - ), + credential_values=incoming_values, + credential_values_to_delete=credential.credential_values_to_delete, ) merged_credential: Final = update_db_credential(db_credential, patch) credential_object_jsonified: Final = cast( # cast-ok: deep-copies a model_dump, so keys are str - "dict[str, object]", jsonify_object(merged_credential.model_dump()) + "dict[str, object]", jsonify_object(merged_credential.model_dump(exclude_none=True)) ) await credentials_repository.update_by_name( credential_name, @@ -371,29 +555,7 @@ async def update_credential( ) # Sync in-memory credential_list (skip if not in memory - e.g., proxy restarted) - new_name: Final = merged_credential.credential_name - existing_in_memory: CredentialItem | None = None - for cred in litellm.credential_list: - if cred.credential_name == credential_name: - existing_in_memory = cred - break - - if existing_in_memory is not None: - in_memory_values: Final = dict(existing_in_memory.credential_values or {}) - if patch.credential_values: - in_memory_values.update(patch.credential_values) - in_memory_info: Final = dict(existing_in_memory.credential_info or {}) - if patch.credential_info: - in_memory_info.update(patch.credential_info) - updated_in_memory: Final = CredentialItem( - credential_name=new_name, - credential_values=in_memory_values, - credential_info=in_memory_info, - ) - # Remove old entry if renamed, then use upsert_credentials to handle duplicates - if new_name != credential_name: - litellm.credential_list = [c for c in litellm.credential_list if c.credential_name != credential_name] - CredentialAccessor.upsert_credentials([updated_in_memory]) + _sync_in_memory_credential(patch, credential_name, merged_credential.credential_name) return {"success": True, "message": "Credential updated successfully"} except Exception as e: diff --git a/litellm/proxy/health_check.py b/litellm/proxy/health_check.py index d5d0123cb6a..468dc7f19bb 100644 --- a/litellm/proxy/health_check.py +++ b/litellm/proxy/health_check.py @@ -32,11 +32,17 @@ from litellm.router_utils.auto_router_model_naming import ( classify_strategy_router_model, strategy_router_dependencies, ) +from litellm.types.utils import secret_bearing_wif_litellm_params, server_owned_wif_litellm_params -# Provider routing fields. Allowed for proxy admins so they can see which -# region/version a deployment is checking; gated at the endpoint layer for -# non-admin callers (see _strip_admin_only_fields_from_health_result). -ADMIN_ONLY_HEALTH_DISPLAY_PARAMS: Final = ("api_base", "api_version", "aws_bedrock_runtime_endpoint") +# Provider routing and workload identity federation fields. Allowed for proxy admins so they can +# see which region/version a deployment is checking and which identity it federates as; gated at +# the endpoint layer for non-admin callers (see _strip_admin_only_fields_from_health_result). +ADMIN_ONLY_HEALTH_DISPLAY_PARAMS: Final = ( + "api_base", + "api_version", + "aws_bedrock_runtime_endpoint", + *(name for name in server_owned_wif_litellm_params if name not in secret_bearing_wif_litellm_params), +) MINIMAL_DISPLAY_PARAMS: Final = frozenset({"model", "mode_error"}) diff --git a/litellm/proxy/health_endpoints/_health_endpoints.py b/litellm/proxy/health_endpoints/_health_endpoints.py index 8804a190d4d..00f0ef1756f 100644 --- a/litellm/proxy/health_endpoints/_health_endpoints.py +++ b/litellm/proxy/health_endpoints/_health_endpoints.py @@ -42,6 +42,7 @@ from litellm.proxy.auth.auth_checks import ( ) from litellm.proxy.auth.auth_utils import ( _BANNED_REQUEST_BODY_PARAMS, # pyright: ignore[reportPrivateUsage] # one canonical list, shared with the request-body check + reject_server_owned_wif_params, ) from litellm.proxy.auth.model_checks import get_key_models from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -956,9 +957,9 @@ def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool: def _strip_admin_only_fields_from_health_result(result: dict) -> dict: """ - Return a copy of the /health response with provider routing fields - (``ADMIN_ONLY_HEALTH_DISPLAY_PARAMS``) removed from each healthy/unhealthy - endpoint entry. Used to hide those fields from non-admin callers while + Return a copy of the /health response with the admin-only fields (provider routing plus the + workload identity federation params naming the identity a deployment mints as) removed from + each healthy/unhealthy endpoint entry. Used to hide those fields from non-admin callers while still showing them which deployments they own and whether each one is healthy. Proxy admins receive the unmodified result. """ @@ -2203,6 +2204,7 @@ async def test_model_connection( "Could not find model %s in router: %s. Proceeding with request params only.", model_name, e ) + reject_server_owned_wif_params(request_litellm_params) # Merge: config params (from proxy config) as base, request params override litellm_params = { **_config_base_for_health_check( @@ -2228,12 +2230,16 @@ async def test_model_connection( await ModelManagementAuthChecks.can_user_make_model_call( model_params=Deployment( model_name="test_model", - litellm_params=LiteLLM_Params(**litellm_params), + litellm_params=LiteLLM_Params.model_validate(litellm_params), model_info=resolved_model_info, ), user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, premium_user=premium_user, + # The probe is a write of the caller's own params onto the stored deployment, so the + # caller's params are the incoming side: a probe that redirects a federated + # deployment's api_base is an admin's action, an unmodified probe of it is not. + incoming_params=request_litellm_params, ) raw_params_mode: Final[object] = litellm_params.pop("mode", None) probe_mode: Final = ( @@ -2260,7 +2266,7 @@ async def test_model_connection( "result": cleaned_result, } - except HTTPException as e: + except (HTTPException, ProxyException) as e: raise e except Exception as e: verbose_proxy_logger.debug("litellm.proxy.health_endpoints.test_model_connection(): Exception occurred - %s", e) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 50bb831d169..11a0075ffd0 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -63,6 +63,12 @@ from litellm.proxy.common_utils.config_sync_pubsub import ( coordination_redis_cache, publish_config_change, ) +from litellm.proxy.common_utils.credential_hydration import ( + effective_server_owned_wif_fields, + hydrate_named_credential, + submitted_litellm_params, + write_touches_federation_surface, +) from litellm.proxy.common_utils.encrypt_decrypt_utils import ( decrypt_value_helper, encrypt_value_helper, @@ -362,6 +368,27 @@ def _raise_on_strategy_router_write_violation( ) +def _reject_non_admin_blocked_flag_on_create( + blocked: bool | None, + user_api_key_dict: UserAPIKeyAuth, +) -> None: + """Same proxy-admin-only rule patch_model applies to the blocked flag: a team admin passed + the team-scoped auth check above, but must not be able to create a model already paused out + from under the proxy admin. + + Only a blocking value is refused. A create that sends ``blocked: false`` asks for the state + every create already lands in, and dashboards and SDKs send the whole model shape on every + create, so refusing the flag's presence would turn a working non-admin create into a 403. + """ + if blocked and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN: + raise ProxyException( + message="Only proxy admins can set a model's blocked flag.", + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param="blocked", + ) + + def _stored_credential_name(existing_litellm_params: GenericLiteLLMParams | None) -> str | None: if existing_litellm_params is None or existing_litellm_params.litellm_credential_name is None: return None @@ -437,6 +464,20 @@ def _effective_complexity_router_config( ).effective +def _decrypted_litellm_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object]: + dumped: Final[Mapping[str, object]] = litellm_params.model_dump(exclude_none=True) + return MappingProxyType( + { + name: ( + decrypt_value_helper(value=value, key=name, exception_type="debug", return_original_value=True) + if isinstance(value, str) + else value + ) + for name, value in dumped.items() + } + ) + + def _effective_model( incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None ) -> str | None: @@ -1161,6 +1202,7 @@ async def patch_model( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, premium_user=premium_user, + incoming_params=submitted_litellm_params(patch_data.litellm_params), member_operation="update", incoming_model_params=patch_data, ) @@ -1524,6 +1566,8 @@ async def _add_model_to_db( } if model_params.model_info.id is not None: _data["model_id"] = model_params.model_info.id + if model_params.blocked is not None: + _data["blocked"] = model_params.blocked _create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above if not should_create_model_in_db: return LiteLLM_ProxyModelTable(**_data) @@ -2105,12 +2149,51 @@ class ModelManagementAuthChecks: ) return True + @staticmethod + async def _reject_non_admin_wif_write( + *, + model_params: Deployment, + incoming_params: Mapping[str, object] | None, + user_api_key_dict: UserAPIKeyAuth, + prisma_client: PrismaClient, + ) -> None: + if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN: + return + if not write_touches_federation_surface(incoming_params): + return + stored: Final = _decrypted_litellm_params(model_params.litellm_params) + wif_fields: Final = await effective_server_owned_wif_fields(stored, incoming_params, prisma_client) + if wif_fields: + # ProxyException rather than HTTPException so the offending field stays a structured + # `param`, which is the contract the narrower gate this replaced already published. + raise ProxyException( + message=( + f"Only proxy admins can change the credentials of a deployment configured for " + f"workload identity federation ({wif_fields[0]!r})." + ), + type=ProxyErrorTypes.auth_error.value, + code=status.HTTP_403_FORBIDDEN, + param=wif_fields[0], + ) + # A name the caller expects an admin to create later would resolve to nothing today and + # start federating the moment it exists, so a non-admin may only attach one that is already there. + named: Final = None if incoming_params is None else incoming_params.get("litellm_credential_name") + if isinstance(named, str) and await hydrate_named_credential(named, prisma_client) is None: + raise ProxyException( + message=f"No credential named {named!r} exists.", + type=ProxyErrorTypes.bad_request_error.value, + code=status.HTTP_400_BAD_REQUEST, + param="litellm_credential_name", + ) + @staticmethod async def can_user_make_model_call( model_params: Deployment, user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, premium_user: bool, + *, + incoming_params: Mapping[str, object] | None, allow_missing_team: bool = False, member_operation: Literal["create", "update"] | None = None, incoming_model_params: updateDeployment | None = None, @@ -2120,6 +2203,19 @@ class ModelManagementAuthChecks: LitellmUserRoles.INTERNAL_USER_VIEW_ONLY, ): raise HTTPException(status_code=403, detail="View-only users cannot manage models.") + # Federation fields choose which server-side secret is read and where the org-scoped token + # it buys is sent, so only a proxy admin may point a federated deployment somewhere else. + # Evaluated on the RESULTING deployment: a patch attaching a credential by name inherits + # whatever that credential holds. `incoming_params` carries only the fields the write set + # and is keyword-only with no default, so a new write path cannot typecheck without + # deciding what it writes. + await ModelManagementAuthChecks._reject_non_admin_wif_write( + model_params=model_params, + incoming_params=incoming_params, + user_api_key_dict=user_api_key_dict, + prisma_client=prisma_client, + ) + ## Check team model auth if model_params.model_info.team_id is not None: team_obj_row: Final = await _repo_team_table(prisma_client).find_unique( @@ -2233,6 +2329,7 @@ async def delete_model( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, premium_user=premium_user, + incoming_params=None, allow_missing_team=True, ) @@ -2355,7 +2452,6 @@ async def delete_team_model_alias( return removed_model_aliases -#### [BETA] - This is a beta endpoint, format might change based on user feedback. - https://github.com/BerriAI/litellm/issues/964 @router.post( "/model/new", description="Allows adding new models to the model list in the config.yaml", @@ -2424,10 +2520,13 @@ async def add_new_model( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, premium_user=premium_user, + incoming_params=submitted_litellm_params(model_params.litellm_params), member_operation="create", ) member_write: Final = write_authorization if isinstance(write_authorization, MemberAutoRouterWrite) else None + _reject_non_admin_blocked_flag_on_create(model_params.blocked, user_api_key_dict) + ModelManagementAuthChecks.can_user_attach_credential( litellm_params=model_params.litellm_params, user_api_key_dict=user_api_key_dict, @@ -2625,6 +2724,7 @@ async def update_model( user_api_key_dict=user_api_key_dict, prisma_client=prisma_client, premium_user=premium_user, + incoming_params=submitted_litellm_params(model_params.litellm_params), member_operation="update", incoming_model_params=model_params, ) diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index 923a6cc5743..b712065b345 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -45,7 +45,7 @@ from litellm.constants import ( BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES, ) from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix -from litellm.llms.anthropic.common_utils import AnthropicModelInfo +from litellm.llms.anthropic.common_utils import AnthropicModelInfo, merge_anthropic_beta_headers from litellm.llms.azure.passthrough.transformation import ( foreign_azure_deployment, is_azure_body_model_inference_endpoint, @@ -914,7 +914,9 @@ async def anthropic_proxy_route( is_streaming_request: Final = await is_streaming_request_fn(request) ## CREATE PASS-THROUGH - auth_header: Final = AnthropicModelInfo.get_auth_header(anthropic_api_key or None) + auth_header: Final = await AnthropicModelInfo.aget_auth_header( + anthropic_api_key or None, allow_workload_identity=True + ) endpoint_func: Final = create_pass_through_route( endpoint=endpoint, target=str(updated_url), @@ -2645,9 +2647,20 @@ def _upstream_headers_for_anthropic_route( caller_headers: Final = _caller_headers_without_litellm_secrets( request, user_api_key_dict, _HEADERS_NEVER_FORWARDED_TO_ANTHROPIC ) - if proxy_auth_header is None and _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS.isdisjoint(caller_headers): - raise HTTPException(status_code=401, detail=_CREDENTIALLESS_ANTHROPIC_MISSING_CREDENTIAL_DETAIL) - return MappingProxyType({**caller_headers, **(proxy_auth_header or {})}) + if proxy_auth_header is None: + if _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS.isdisjoint(caller_headers): + raise HTTPException(status_code=401, detail=_CREDENTIALLESS_ANTHROPIC_MISSING_CREDENTIAL_DETAIL) + return caller_headers + forwarded: Final = MappingProxyType( + {name: value for name, value in caller_headers.items() if name not in _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS} + ) + caller_beta, credential_beta = caller_headers.get("anthropic-beta"), proxy_auth_header.get("anthropic-beta") + merged_beta: Final = ( + {"anthropic-beta": merge_anthropic_beta_headers(caller_beta, credential_beta)} + if caller_beta and credential_beta + else {} + ) + return MappingProxyType({**forwarded, **proxy_auth_header, **merged_beta}) def _upstream_headers_for_bedrock_agent_runtime_route( diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 89f19a16da4..f40e541dd86 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -413,6 +413,7 @@ from litellm.proxy.common_utils.callback_utils import initialize_callbacks_on_pr from litellm.proxy.common_utils.codex_model_catalog import codex_model_list_body from litellm.proxy.common_utils.config_includes import resolve_include_file_path, resolve_includes from litellm.proxy.common_utils.config_sync_pubsub import ConfigSyncSubscriber +from litellm.proxy.common_utils.credential_hydration import decrypted_or_stored from litellm.proxy.common_utils.debug_utils import init_verbose_loggers from litellm.proxy.common_utils.debug_utils import router as debugging_endpoints_router from litellm.proxy.common_utils.discoverable_model_filter import discoverable_rows, undiscoverable_model_names @@ -901,6 +902,7 @@ from litellm.types.router import ( RoutingGroup, RoutingPlugin, SearchToolTypedDict, + holds_secret_pointer, updateDeployment, ) from litellm.types.router import ModelInfo as RouterModelInfo @@ -5793,11 +5795,11 @@ class ProxyConfig: return config return { - key: self._resolved_config_value(value=value, depth=depth, max_depth=max_depth) + key: self._resolved_config_value(key=key, value=value, depth=depth, max_depth=max_depth) for key, value in config.items() } - def _resolved_config_value(self, value: object, depth: int, max_depth: int) -> object: + def _resolved_config_value(self, key: str, value: object, depth: int, max_depth: int) -> object: if isinstance(value, dict): return self._check_for_os_environ_vars(config=value, depth=depth + 1, max_depth=max_depth) if isinstance(value, list): @@ -5807,7 +5809,7 @@ class ProxyConfig: else item for item in value ] - if isinstance(value, str) and value.startswith("os.environ/"): + if isinstance(value, str) and value.startswith("os.environ/") and not holds_secret_pointer(key): resolved: Final = get_secret(value) if resolved is None and secret_manager_would_be_consulted(value): verbose_proxy_logger.warning("%s is absent from the configured secret manager", value) @@ -6938,7 +6940,7 @@ class ProxyConfig: for model in model_list: ### LOAD FROM os.environ/ ### for k, v in model["litellm_params"].items(): - if isinstance(v, str) and v.startswith("os.environ/"): + if isinstance(v, str) and v.startswith("os.environ/") and not holds_secret_pointer(k): model["litellm_params"][k] = get_secret(v) validate_deployment_max_agentic_loops(model) validate_deployment_complexity_router_placement(model) @@ -7349,7 +7351,7 @@ class ProxyConfig: for model in model_list: ### LOAD FROM os.environ/ ### for k, v in model["litellm_params"].items(): - if isinstance(v, str) and v.startswith("os.environ/"): + if isinstance(v, str) and v.startswith("os.environ/") and not holds_secret_pointer(k): model["litellm_params"][k] = get_secret(v) ## check if they have model-id's ## @@ -7394,7 +7396,11 @@ class ProxyConfig: return value decrypted_value: Final = decrypt_value_helper(value=value, key=key, return_original_value=True) - if isinstance(decrypted_value, str) and decrypted_value.startswith("os.environ/"): + if ( + isinstance(decrypted_value, str) + and decrypted_value.startswith("os.environ/") + and not holds_secret_pointer(key) + ): return get_secret(decrypted_value) return decrypted_value @@ -9118,7 +9124,7 @@ class ProxyConfig: decrypted_credential_values: Final = {} for k, v in credential_object.credential_values.items(): - decrypted_credential_values[k] = decrypt_value_helper(value=v, key=k) or v + decrypted_credential_values[k] = decrypted_or_stored(k, v) credential_object.credential_values = decrypted_credential_values return credential_object diff --git a/litellm/router.py b/litellm/router.py index 54fffae8dbd..0b9f12c8da3 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -307,6 +307,7 @@ from litellm.types.router import ( RoutingStrategy, SearchToolTypedDict, TaggedPreRoutingStrategy, + holds_secret_pointer, ) from litellm.types.services import ServiceTypes from litellm.types.utils import ( @@ -3987,7 +3988,7 @@ class Router: model_info["original_model_id"] = original_model_id deployment_pydantic_obj: Final = Deployment( model_name=model_group, - litellm_params=LiteLLM_Params(**dynamic_litellm_params), + litellm_params=LiteLLM_Params.model_validate(dynamic_litellm_params), model_info=model_info, ) Router._register_deployment_pricing(deployment=deployment_pydantic_obj) @@ -9000,13 +9001,10 @@ class Router: if access_windows_error is not None: raise ValueError(access_windows_error) zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None - litellm_params: Final[LiteLLM_Params] = LiteLLM_Params( - **( # pyright: ignore[reportArgumentType] # untyped merged dict; already true for every field here - _litellm_params - if zeroed_pricing is None - else MappingProxyType({**_litellm_params, **zeroed_pricing}) - ) + merged_params: Final[Mapping[str, Any]] = ( + _litellm_params if zeroed_pricing is None else MappingProxyType({**_litellm_params, **zeroed_pricing}) ) + litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(**merged_params) warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params) deployment = Deployment( **deployment_info, @@ -9362,7 +9360,7 @@ class Router: continue deployment = Deployment( model_name=model_name, - litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params(**lp)), + litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params.model_validate(lp)), model_info=(entry.get("model_info") if isinstance(entry, dict) else entry.model_info), ) if self._has_registered_strategy(self.adaptive_routers, model_name, self._deployment_tags(deployment)): @@ -9584,7 +9582,7 @@ class Router: ## check if litellm params in os.environ if isinstance(_litellm_params, dict): for k, v in _litellm_params.items(): - if isinstance(v, str) and v.startswith("os.environ/"): + if isinstance(v, str) and v.startswith("os.environ/") and not holds_secret_pointer(k): _litellm_params[k] = get_secret(v) _model_info: dict = model.pop("model_info", {}) @@ -10791,7 +10789,7 @@ class Router: if isinstance(litellm_params_data, LiteLLM_Params): litellm_params = litellm_params_data elif isinstance(litellm_params_data, dict) and "model" in litellm_params_data: - litellm_params = LiteLLM_Params(**litellm_params_data) + litellm_params = LiteLLM_Params.model_validate(litellm_params_data) else: raise ValueError( f"Deployment missing valid litellm_params. " @@ -12633,7 +12631,7 @@ class Router: if allowed_model_region is not None: if not is_region_allowed( - litellm_params=LiteLLM_Params(**_litellm_params), + litellm_params=LiteLLM_Params.model_validate(_litellm_params), allowed_model_region=allowed_model_region, ): invalid_model_indices.add(idx) @@ -12651,7 +12649,7 @@ class Router: _, ) = litellm.get_llm_provider( model=_dep_model_for_params, - litellm_params=LiteLLM_Params(**_litellm_params), + litellm_params=LiteLLM_Params.model_validate(_litellm_params), ) except Exception as e: # noqa: BLE001 # best-effort filter: an unresolvable provider must not fail the request verbose_router_logger.debug( diff --git a/litellm/router_utils/clientside_credential_handler.py b/litellm/router_utils/clientside_credential_handler.py index 55b246b22a7..186772925a9 100644 --- a/litellm/router_utils/clientside_credential_handler.py +++ b/litellm/router_utils/clientside_credential_handler.py @@ -13,8 +13,16 @@ Ensures cooldowns are applied correctly. from typing import Final +from litellm.types.utils import server_owned_wif_litellm_params + clientside_credential_keys: Final = ["api_key", "api_base", "base_url"] +# Set on a deployment whose api_base was client-redirected, so the Anthropic auth path refuses to +# mint a federation token there even when WIF is configured only through ANTHROPIC_* env vars (which +# cannot be cleared from litellm_params). +DISABLE_WORKLOAD_IDENTITY_PARAM: Final = "anthropic_disable_workload_identity_federation" +_WIF_CLEAR_ON_BASE_OVERRIDE: Final = tuple(sorted(server_owned_wif_litellm_params)) + def _admin_config_fields_to_clear_on_base_override() -> list[str]: """ @@ -59,6 +67,14 @@ def _admin_config_fields_to_clear_on_base_override() -> list[str]: # ``api_base`` for the same reason as the OCI entries above. "nvcf_function_id", "use_ssl", + # Workload-identity federation minting fields, restated here from + # server_owned_wif_litellm_params the same way azure_ad_token above is restated + # despite also being declared on CredentialLiteLLMParams (hence covered by + # typed_fields too): a federation token minted for a client-redirected api_base + # would send the workload's OIDC assertion, and then the minted bearer, to the + # caller-chosen host, so this list must stay correct even if a field is ever + # dropped from the typed model. + *_WIF_CLEAR_ON_BASE_OVERRIDE, ] return typed_fields + kwargs_only_fields @@ -101,5 +117,6 @@ def get_dynamic_litellm_params(litellm_params: dict, request_kwargs: dict) -> di litellm_params.pop(field, None) if field in request_kwargs: litellm_params[field] = request_kwargs[field] + litellm_params[DISABLE_WORKLOAD_IDENTITY_PARAM] = True return litellm_params diff --git a/litellm/router_utils/fallback_event_handlers.py b/litellm/router_utils/fallback_event_handlers.py index 3142fd5fb98..74504d76736 100644 --- a/litellm/router_utils/fallback_event_handlers.py +++ b/litellm/router_utils/fallback_event_handlers.py @@ -28,7 +28,7 @@ from litellm.router_utils.cooldown_handlers import ( from litellm.router_utils.router_callbacks.track_deployment_metrics import ( increment_deployment_failures_for_current_minute, ) -from litellm.types.router import LiteLLMParamsTypedDict +from litellm.types.router import LiteLLMParamsTypedDict, reject_server_owned_wif_params if TYPE_CHECKING: from litellm.router import Router as _Router @@ -640,6 +640,14 @@ async def run_async_fallback( failed_model_group: Final = get_pre_routing_selection(kwargs) or original_model_group attempted.record(failed_model_group) + # A dict target is merged straight into kwargs below, and kwargs win over the deployment's own + # params, so a stored key/team/global fallback could otherwise set a federation field that the + # request itself is forbidden to carry. Checked here rather than at the merge: inside the loop + # the refusal would be caught as a per-target failure and quietly skipped to the next one. + for target in fallback_model_group: + if isinstance(target, dict): + reject_server_owned_wif_params(target) + for mg in fallback_model_group: if mg == failed_model_group: continue diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index f09ddd1d5a9..fbed4fb8e75 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -49,6 +49,10 @@ def _oidc_token_cache_ttl(oidc_token: str, max_ttl: int) -> int: _DEFAULT_OIDC_ALLOWED_CREDENTIAL_DIRS: Final = ("/var/run/secrets", "/run/secrets") +class OidcPathNotAllowedError(ValueError): + """An ``oidc/file/`` path was rejected by the credential-directory allowlist.""" + + def _get_oidc_allowed_credential_dirs() -> list[str]: """ Return the absolute, normalized list of directories from which @@ -73,7 +77,7 @@ def _resolve_oidc_file_path(requested_path: str) -> str: credential directories. Raises ``ValueError`` otherwise. """ if not os.path.isabs(requested_path): - raise ValueError( + raise OidcPathNotAllowedError( "oidc/file path must be absolute. Use the format " "'oidc/file//var/run/secrets/' (note the leading slash " "after 'oidc/file/')." @@ -87,7 +91,7 @@ def _resolve_oidc_file_path(requested_path: str) -> str: # commonpath raises when paths are on different drives (Windows); # treat as not-matching and continue. continue - raise ValueError( + raise OidcPathNotAllowedError( "oidc/file path is outside the allowed credential directories. " "Set LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS to extend the allowlist." ) diff --git a/litellm/types/litellm_params.py b/litellm/types/litellm_params.py index 51f6671e9d6..9119fa16b5d 100644 --- a/litellm/types/litellm_params.py +++ b/litellm/types/litellm_params.py @@ -104,10 +104,41 @@ class BedrockBatchConnection: bedrock_tags: Sequence[Mapping[str, str]] | None = None +@dataclass(frozen=True, slots=True, kw_only=True) +class AnthropicFederationConnection: + anthropic_federation_rule_id: str | None = None + anthropic_organization_id: str | None = None + anthropic_service_account_id: str | None = None + anthropic_federation_workspace_id: str | None = None + anthropic_identity_token_file: str | None = None + anthropic_identity_token: str | None = None + anthropic_identity_source: str | None = None + anthropic_issuer_url: str | None = None + anthropic_issuer_subject: str | None = None + anthropic_issuer_audience: str | None = None + anthropic_issuer_ttl_seconds: int | None = None + anthropic_issuer_signing_key_ref: str | None = None + anthropic_keycloak_token_url: str | None = None + anthropic_keycloak_client_id: str | None = None + anthropic_keycloak_auth_method: str | None = None + anthropic_keycloak_client_secret_ref: str | None = None + anthropic_keycloak_scope: str | None = None + anthropic_disable_workload_identity_federation: bool | None = None + + +@dataclass(frozen=True, slots=True, kw_only=True) +class OpenAIFederationConnection: + openai_identity_provider_id: str | None = None + openai_service_account_id: str | None = None + openai_identity_token_file: str | None = None + + @dataclass(frozen=True, slots=True, kw_only=True) class ConnectionSettings: provider: ProviderConnection bedrock_batch: BedrockBatchConnection + anthropic_federation: AnthropicFederationConnection + openai_federation: OpenAIFederationConnection @dataclass(frozen=True, slots=True, kw_only=True) diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index ee357cd6581..46693009bfd 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -784,5 +784,6 @@ ANTHROPIC_FINE_GRAINED_TOOL_STREAMING_BETA_HEADER: Final = "fine-grained-tool-st # OAuth constants ANTHROPIC_OAUTH_TOKEN_PREFIX: Final = "sk-ant-oat" ANTHROPIC_OAUTH_BETA_HEADER: Final = "oauth-2025-04-20" +ANTHROPIC_TOKEN_EXCHANGE_PATH: Final = "/v1/oauth/token" ANTHROPIC_PROMPT_CACHING_SCOPE_BETA_HEADER: Final = "prompt-caching-scope-2026-01-05" diff --git a/litellm/types/router.py b/litellm/types/router.py index 2ab1a1185ed..e131a7184d4 100644 --- a/litellm/types/router.py +++ b/litellm/types/router.py @@ -4,7 +4,7 @@ litellm.Router Types - includes RouterConfig, UpdateRouterConfig, ModelInfo etc import datetime import enum -from collections.abc import Mapping, Sequence +from collections.abc import Container, Mapping, Sequence from dataclasses import dataclass from typing import TYPE_CHECKING, Annotated, Any, ClassVar, Final, Generic, Literal, TypeVar, get_type_hints from zoneinfo import ZoneInfo, ZoneInfoNotFoundError @@ -34,6 +34,10 @@ from .utils import ( ModelResponse, StandardLoggingRoutingDecision, ) +from .utils import ( + # private alias: `from .types.router import *` would rebind a public Final in litellm/__init__.py + server_owned_wif_litellm_params as _server_owned_wif_litellm_params, +) class ConfigurableClientsideParamsCustomAuth(TypedDict): @@ -374,6 +378,69 @@ class CredentialLiteLLMParams(BaseModel): ## IBM WATSONX ## watsonx_region_name: str | None = None + ## ANTHROPIC WORKLOAD IDENTITY FEDERATION ## + # Without these, get_deployment_credentials_with_provider silently drops a + # litellm_params-configured WIF setup before files/batches/passthrough callers see + # it, the same #30235-shaped gap azure_ad_token above was added to close. + anthropic_federation_rule_id: str | None = None + anthropic_organization_id: str | None = None + anthropic_service_account_id: str | None = None + anthropic_federation_workspace_id: str | None = None + anthropic_identity_token_file: str | None = None + anthropic_identity_token: str | None = None + anthropic_identity_source: str | None = None + anthropic_issuer_url: str | None = None + anthropic_issuer_subject: str | None = None + anthropic_issuer_audience: str | None = None + anthropic_issuer_ttl_seconds: int | None = None + anthropic_issuer_signing_key_ref: str | None = None + anthropic_keycloak_token_url: str | None = None + anthropic_keycloak_client_id: str | None = None + anthropic_keycloak_auth_method: str | None = None + anthropic_keycloak_client_secret_ref: str | None = None + anthropic_keycloak_scope: str | None = None + # Server-set when a client redirects api_base. Declared so it survives the strict dump the + # other federation fields above are declared for, rather than being rebuilt away in transit. + anthropic_disable_workload_identity_federation: bool | None = None + + ## OPENAI WORKLOAD IDENTITY FEDERATION ## + openai_identity_provider_id: str | None = None + openai_service_account_id: str | None = None + openai_identity_token_file: str | None = None + + +def server_owned_wif_fields_present(fields: Mapping[str, object]) -> tuple[str, ...]: + """Server-owned workload identity federation field names set in ``fields``. + + ``fields`` is a ``litellm_params`` dict (or a credential's ``credential_values`` mapping, + which feeds the same resolution when referenced by name). Derived from + ``server_owned_wif_litellm_params`` rather than hand-copied, so a persistence gate built on + this stays correct when a new WIF field is added there. + """ + return tuple(name for name in _server_owned_wif_litellm_params if fields.get(name) is not None) + + +def server_owned_wif_fields_named(keys: Container[str]) -> tuple[str, ...]: + """Server-owned workload identity federation field names that appear in ``keys``, whatever + value they carry. + + The write gates on credentials need this key-based sibling of ``server_owned_wif_fields_present``: + ``get_litellm_params`` forwards a WIF kwarg on key presence and the federation resolver rejects + a foreign variant's field by key, so a persisted ``{"anthropic_issuer_url": None}`` wedges every + deployment that references the credential even though no value is set. Pass a mapping (its keys + are tested) or a plain collection of key names. + """ + return tuple(name for name in _server_owned_wif_litellm_params if name in keys) + + +_WIF_POINTER_FIELDS: Final = frozenset(name for name in _server_owned_wif_litellm_params if name.endswith("_ref")) + + +def holds_secret_pointer(param_name: str) -> bool: + """A ``*_ref`` federation field is a secret POINTER the identity source dereferences at use + time, so a loader expanding ``os.environ/`` values must leave it as written.""" + return param_name in _WIF_POINTER_FIELDS + _RESERVED_INIT_KEYS: Final = frozenset({"self", "params", "__class__"}) @@ -648,6 +715,9 @@ class Deployment(BaseModel): model_name: str litellm_params: LiteLLM_Params model_info: ModelInfo + # admin-toggled pause flag; mirrors LiteLLM_ProxyModelTable.blocked. None means "don't set it + # on create" -- the Prisma column defaults to False -- rather than "explicitly unblocked". + blocked: bool | None = None model_config = ConfigDict(extra="allow", protected_namespaces=()) @@ -1243,3 +1313,22 @@ class AdaptiveRouterPreferences(BaseModel): quality_tier: int = Field(ge=1, le=3) strengths: list[RequestType] = Field(default_factory=list) + + +def reject_server_owned_wif_params(body: Mapping[str, object]) -> None: + """Raise ``ValueError`` if a mapping that did not come from deployment config carries a + server-owned workload identity federation field. + + These are never settable inline on a client surface, with or without a client-side credential + opt-in. Naming a stored credential that already holds them is the other way in and has its own + gate: ``_check_banned_params`` resolves ``litellm_credential_name`` and refuses a federated one. + This lives here rather than under ``litellm.proxy`` so the router can call it on a + post-authentication merge without core importing from the proxy package. + """ + for param in _server_owned_wif_litellm_params: + if param in body: + raise ValueError( + f"Rejected Request: {param} is a server-owned workload identity federation parameter " + "and cannot be set in a request body. A proxy admin configures it on the deployment " + "or on a stored credential." + ) diff --git a/litellm/types/services.py b/litellm/types/services.py index 00fa9f044cc..8a74be06da6 100644 --- a/litellm/types/services.py +++ b/litellm/types/services.py @@ -25,6 +25,8 @@ class ServiceTypes(str, enum.Enum): AUTH = "auth" PROXY_PRE_CALL = "proxy_pre_call" POD_LOCK_MANAGER = "pod_lock_manager" + ANTHROPIC_WIF = "anthropic_wif" + ANTHROPIC_WIF_CACHE = "anthropic_wif_cache" """ Operational metrics for DB Transaction Queues @@ -67,6 +69,9 @@ DEFAULT_SERVICE_CONFIGS: Final = { ServiceTypes.ROUTER.value: {"metrics": [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM]}, ServiceTypes.AUTH.value: {"metrics": [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM]}, ServiceTypes.PROXY_PRE_CALL.value: {"metrics": [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM]}, + ServiceTypes.ANTHROPIC_WIF.value: {"metrics": [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM]}, + # cache hits are counter-only: no HTTP call happens, so observing a latency would be a lie + ServiceTypes.ANTHROPIC_WIF_CACHE.value: {"metrics": [ServiceMetrics.COUNTER]}, # Operational metrics for DB Transaction Queues ServiceTypes.POD_LOCK_MANAGER.value: {"metrics": [ServiceMetrics.GAUGE]}, ServiceTypes.IN_MEMORY_DAILY_SPEND_UPDATE_QUEUE.value: {"metrics": [ServiceMetrics.GAUGE]}, diff --git a/litellm/types/utils.py b/litellm/types/utils.py index a33eeaccaa3..142f9b14a72 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -55,6 +55,11 @@ from litellm.types.llms.base import ( LiteLLMPydanticObjectBase, ) from litellm.types.mcp import MCPServerCostInfo +from litellm.types.workload_identity import ( + ANTHROPIC_WIF_KWARGS_KEYS, + OPENAI_WIF_KWARGS_KEYS, + WIF_SECRET_BEARING_KEYS, +) from ..litellm_core_utils.core_helpers import map_finish_reason, process_response_headers from . import litellm_params as _litellm_params @@ -3958,6 +3963,11 @@ bedrock_batch_litellm_params: Final = BEDROCK_BATCH_KWARG_NAMES TRUSTED_CALLBACK_VARS_FIELD: Final = _litellm_params.TRUSTED_CALLBACK_VARS_FIELD ADDRESSED_RESPONSE_ID_FIELD: Final = _litellm_params.ADDRESSED_RESPONSE_ID_FIELD +anthropic_wif_litellm_params: Final = tuple(sorted(ANTHROPIC_WIF_KWARGS_KEYS)) +openai_wif_litellm_params: Final = tuple(sorted(OPENAI_WIF_KWARGS_KEYS)) +server_owned_wif_litellm_params: Final = anthropic_wif_litellm_params + openai_wif_litellm_params +secret_bearing_wif_litellm_params: Final = tuple(sorted(WIF_SECRET_BEARING_KEYS)) + all_litellm_params = [ # rebind-ok: two star imports in litellm/__init__.py re-bind it *OWNED_KWARG_NAMES, *KWARG_ARTIFACTS, diff --git a/litellm/types/workload_identity.py b/litellm/types/workload_identity.py new file mode 100644 index 00000000000..b0dbd9566e8 --- /dev/null +++ b/litellm/types/workload_identity.py @@ -0,0 +1,35 @@ +"""litellm_params keys that configure workload identity federation. + +The key sets derive from the federation connection leaves in ``types.litellm_params``, so the +kwargs funnel (``litellm_core_utils.get_litellm_params``) and the request-body ban list +(``types.utils.all_litellm_params``) read one declaration. Every key rides the funnel into +``litellm_params`` and is banned from request bodies, which also covers +``anthropic_disable_workload_identity_federation``: the proxy sets it when a client redirects +``api_base`` so a federated deployment stops minting for a base the caller chose, and a caller +must not be able to set it in either direction. +""" + +from typing import Final + +from litellm.types.litellm_params import AnthropicFederationConnection, OpenAIFederationConnection, wire_names + +ANTHROPIC_WIF_KWARGS_KEYS: Final = frozenset(wire_names(AnthropicFederationConnection)) + +OPENAI_WIF_KWARGS_KEYS: Final = frozenset(wire_names(OpenAIFederationConnection)) + +WIF_SECRET_BEARING_KEYS: Final = frozenset( + { + "anthropic_identity_token", + "anthropic_identity_token_file", + "anthropic_issuer_signing_key_ref", + "anthropic_keycloak_client_secret_ref", + "openai_identity_token_file", + } +) +"""The federation keys whose value is a credential, or the path or reference that reaches one. + +The rest of the sets above name a rule, an organization, a workspace, or a URL: an operator has +to be able to read those back to tell what a deployment federates as. These carry the secret +itself, so no surface displays them to anyone. Splitting the sensitivity out here keeps the +callers that redact them from having to know which provider a field belongs to. +""" diff --git a/litellm/utils.py b/litellm/utils.py index ac2879ca417..a7d7447d7f3 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -7783,9 +7783,8 @@ def _get_valid_models_from_provider_api( if cached_result is not None: return cached_result - models: Final = provider_config.get_models( - api_key=litellm_params.api_key if litellm_params is not None else None, - api_base=litellm_params.api_base if litellm_params is not None else None, + models: Final = provider_config.discover_models( + litellm_params=litellm_params.model_dump(exclude_none=True) if litellm_params is not None else None ) _model_cache.set_cached_model_info(custom_llm_provider, litellm_params, models) diff --git a/terraform/provider/tools/endpointaudit/coverage_allowlist.txt b/terraform/provider/tools/endpointaudit/coverage_allowlist.txt index fe89257e83b..9305d2e109d 100644 --- a/terraform/provider/tools/endpointaudit/coverage_allowlist.txt +++ b/terraform/provider/tools/endpointaudit/coverage_allowlist.txt @@ -65,6 +65,7 @@ GET /user/spend/report # Admin UI helper endpoints; serve UI forms and caller-scoped views, not desired state GET /budget/settings +GET /credentials/{credential_name}/jwks GET /router/fields GET /guardrails/ui/add_guardrail_settings GET /guardrails/ui/category_yaml/{category_name} diff --git a/tests/unit/batches/test_batch_utils.py b/tests/unit/batches/test_batch_utils.py index b8b922f72a7..8147626fc5e 100644 --- a/tests/unit/batches/test_batch_utils.py +++ b/tests/unit/batches/test_batch_utils.py @@ -1890,6 +1890,40 @@ def test_unparsable_bedrock_batch_usage_warns(caplog): assert "inputTextTokenCount" in caplog.text +class TestFileAccessCredentialsCarryFederation: + """A federated deployment holds no api_key, so the fetch that reads a finished batch's output + has to inherit the federation fields or it cannot authenticate and the batch is never billed.""" + + def test_federation_fields_survive_extraction(self): + from litellm.batches.batch_utils import _extract_file_access_credentials + + credentials = _extract_file_access_credentials( + { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_x", + "anthropic_organization_id": "org-x", + "anthropic_identity_token_file": "/var/run/secrets/anthropic.com/token", + "something_unrelated": "dropped", + } + ) + + assert credentials["anthropic_federation_rule_id"] == "fdrl_x" + assert credentials["anthropic_organization_id"] == "org-x" + assert credentials["anthropic_identity_token_file"] == "/var/run/secrets/anthropic.com/token" + assert "something_unrelated" not in credentials + + def test_every_federation_field_is_carried(self): + """Derived from the kwargs set, so a new federation field is carried without an edit here.""" + from litellm.batches.batch_utils import _extract_file_access_credentials + from litellm.litellm_core_utils.get_litellm_params import ANTHROPIC_WIF_KWARGS_KEYS + + params = {name: f"value-{name}" for name in ANTHROPIC_WIF_KWARGS_KEYS} + + credentials = _extract_file_access_credentials(params) + + assert set(credentials) == set(ANTHROPIC_WIF_KWARGS_KEYS) + + def test_total_cost_bills_cached_tokens_per_line_at_the_batch_cached_rate(): responses_row = _success_row( usage={ diff --git a/tests/unit/integrations/test_prometheus_services.py b/tests/unit/integrations/test_prometheus_services.py index 2303061ede8..a5e95de3ea9 100644 --- a/tests/unit/integrations/test_prometheus_services.py +++ b/tests/unit/integrations/test_prometheus_services.py @@ -135,3 +135,28 @@ def test_services_logger_custom_latency_buckets(): REGISTRY.unregister(collector) except Exception: pass + + +def test_anthropic_wif_services_are_wired_into_the_registry(): + """Reverting the ANTHROPIC_WIF/ANTHROPIC_WIF_CACHE ServiceTypes members or their + DEFAULT_SERVICE_CONFIGS entries must fail here: the exchange service gets counters plus a + latency histogram, while the cache-hit service is counter-only so a hit can never fake a latency.""" + from litellm.types.services import DEFAULT_SERVICE_CONFIGS + + assert ServiceTypes.ANTHROPIC_WIF.value == "anthropic_wif" + assert ServiceTypes.ANTHROPIC_WIF_CACHE.value == "anthropic_wif_cache" + assert DEFAULT_SERVICE_CONFIGS["anthropic_wif"]["metrics"] == [ServiceMetrics.COUNTER, ServiceMetrics.HISTOGRAM] + assert DEFAULT_SERVICE_CONFIGS["anthropic_wif_cache"]["metrics"] == [ServiceMetrics.COUNTER] + + pl = PrometheusServicesLogger() + wif_names = {obj._name for obj in pl.payload_to_prometheus_map["anthropic_wif"]} + assert wif_names == { + "litellm_anthropic_wif_latency", + "litellm_anthropic_wif_failed_requests", + "litellm_anthropic_wif_total_requests", + } + cache_names = {obj._name for obj in pl.payload_to_prometheus_map["anthropic_wif_cache"]} + assert cache_names == { + "litellm_anthropic_wif_cache_failed_requests", + "litellm_anthropic_wif_cache_total_requests", + } diff --git a/tests/unit/litellm_core_utils/test_get_litellm_params.py b/tests/unit/litellm_core_utils/test_get_litellm_params.py index 19a3323ce53..c7c37abd46e 100644 --- a/tests/unit/litellm_core_utils/test_get_litellm_params.py +++ b/tests/unit/litellm_core_utils/test_get_litellm_params.py @@ -20,6 +20,28 @@ from litellm.litellm_core_utils.get_litellm_params import ( ) from litellm.types.litellm_params import ControlOptions + +def _funnel_kwargs_completion_forwards(monkeypatch, wif_kwargs: dict[str, object]) -> dict[str, object]: + """completion() names its get_litellm_params arguments one by one, so a key the funnel knows + is still dropped unless that call site forwards it from its own kwargs.""" + from unittest.mock import MagicMock + + import litellm + import litellm.main as litellm_main + + spy = MagicMock(wraps=litellm_main.get_litellm_params) + monkeypatch.setattr( # test-quality-ok: completion() has no injection seam for its kwargs funnel + litellm_main, "get_litellm_params", spy + ) + litellm.completion( + model="anthropic/claude-sonnet-5", + messages=[{"role": "user", "content": "hi"}], + mock_response="ok", + **wif_kwargs, + ) + return spy.call_args.kwargs + + NAMED_PRICE_PARAMS: Final = frozenset( { "input_cost_per_token", @@ -357,3 +379,126 @@ def test_drop_params_strings_reach_litellm_params_as_flags( value: str | bool | None, expected: bool | None ) -> None: assert get_litellm_params(drop_params=value)["drop_params"] is expected + + +class TestAnthropicWifKeys: + """The six anthropic_* WIF keys need dual registration: carried by the kwargs + funnel into litellm_params (where the Anthropic auth tier reads them) AND + listed in all_litellm_params (so the extra_body sweep never sends them to + /v1/messages).""" + + SIX_KEYS = { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_service_account_id": "svcacct_1", + "anthropic_federation_workspace_id": "wrkspc_1", + "anthropic_identity_token_file": "/var/run/secrets/tok", + "anthropic_identity_token": "oidc/env/TOK", + } + + def test_keys_survive_into_litellm_params(self): + params = get_litellm_params(**self.SIX_KEYS) + for key, value in self.SIX_KEYS.items(): + assert params[key] == value + + def test_keys_are_forwarded_from_completion_kwargs(self, monkeypatch): + forwarded = _funnel_kwargs_completion_forwards(monkeypatch, self.SIX_KEYS) + assert {key: forwarded[key] for key in self.SIX_KEYS} == self.SIX_KEYS + + def test_keys_stay_out_of_the_provider_body(self): + from litellm.types.utils import all_litellm_params + + for key in self.SIX_KEYS: + assert key in all_litellm_params + + def test_keys_absent_when_not_configured(self): + params = get_litellm_params() + for key in self.SIX_KEYS: + assert key not in params + + +class TestAnthropicWifIdentitySourceKeys: + """Phase 1 adds 11 more anthropic_* WIF keys (the anthropic_identity_source discriminator + plus the internal_issuer/keycloak identity-source fields) that need the same dual + registration as the original six tested above.""" + + NEW_KEYS = { + "anthropic_identity_source": "keycloak", + "anthropic_issuer_url": "https://issuer.example", + "anthropic_issuer_subject": "svc-account", + "anthropic_issuer_audience": "https://api.anthropic.com", + "anthropic_issuer_ttl_seconds": "300", + "anthropic_issuer_signing_key_ref": "oidc/env/ISSUER_KEY", + "anthropic_keycloak_token_url": "https://kc.example/realms/r/protocol/openid-connect/token", + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_auth_method": "client_secret_basic", + "anthropic_keycloak_client_secret_ref": "oidc/env/KC_SECRET", + "anthropic_keycloak_scope": "anthropic-wif", + # Server-set when a client redirects api_base; carried here so it is not dropped in transit + "anthropic_disable_workload_identity_federation": True, + } + + def test_new_keys_are_exactly_the_non_legacy_registered_set(self): + """Fails the moment a key is added to ANTHROPIC_WIF_KWARGS_KEYS without a matching entry + here (or vice versa), catching drift between what wif.py dispatches on and what this + test (and the funnel/provider-body tests below) actually exercises.""" + from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS + + assert set(self.NEW_KEYS) == ANTHROPIC_WIF_KWARGS_KEYS - set(TestAnthropicWifKeys.SIX_KEYS) + + def test_keys_survive_into_litellm_params(self): + params = get_litellm_params(**self.NEW_KEYS) + for key, value in self.NEW_KEYS.items(): + assert params[key] == value + + def test_keys_are_forwarded_from_completion_kwargs(self, monkeypatch): + forwarded = _funnel_kwargs_completion_forwards(monkeypatch, self.NEW_KEYS) + assert {key: forwarded[key] for key in self.NEW_KEYS} == self.NEW_KEYS + + def test_keys_stay_out_of_the_provider_body(self): + from litellm.types.utils import all_litellm_params + + for key in self.NEW_KEYS: + assert key in all_litellm_params + + def test_keys_absent_when_not_configured(self): + params = get_litellm_params() + for key in self.NEW_KEYS: + assert key not in params + + +class TestOpenAIWifKeys: + """The three openai_* WIF keys carry a deployment's federation identity through the kwargs + funnel into litellm_params (where the OpenAI client factory reads them) and stay out of the + provider body, exactly like the anthropic_* keys above.""" + + THREE_KEYS = { + "openai_identity_provider_id": "idp_1", + "openai_service_account_id": "user-1", + "openai_identity_token_file": "/var/run/secrets/tokens/openai", + } + + def test_keys_are_exactly_the_registered_set(self): + from litellm.types.workload_identity import OPENAI_WIF_KWARGS_KEYS + + assert set(self.THREE_KEYS) == OPENAI_WIF_KWARGS_KEYS + + def test_keys_survive_into_litellm_params(self): + params = get_litellm_params(**self.THREE_KEYS) + for key, value in self.THREE_KEYS.items(): + assert params[key] == value + + def test_keys_are_forwarded_from_completion_kwargs(self, monkeypatch): + forwarded = _funnel_kwargs_completion_forwards(monkeypatch, self.THREE_KEYS) + assert {key: forwarded[key] for key in self.THREE_KEYS} == self.THREE_KEYS + + def test_keys_stay_out_of_the_provider_body(self): + from litellm.types.utils import all_litellm_params + + for key in self.THREE_KEYS: + assert key in all_litellm_params + + def test_keys_absent_when_not_configured(self): + params = get_litellm_params() + for key in self.THREE_KEYS: + assert key not in params diff --git a/tests/unit/llms/anthropic/batches/test_handler.py b/tests/unit/llms/anthropic/batches/test_handler.py index 28b84123482..9b9bddae678 100644 --- a/tests/unit/llms/anthropic/batches/test_handler.py +++ b/tests/unit/llms/anthropic/batches/test_handler.py @@ -13,6 +13,8 @@ The sync ``retrieve_batch`` dispatch (``_is_async`` true -> coroutine, false -> asyncio.run) is exercised directly. """ +import asyncio +import threading from unittest.mock import AsyncMock, MagicMock, patch import httpx @@ -34,9 +36,7 @@ def _ok_batch_response(): "ended_at": "2024-09-24T11:00:00Z", "request_counts": {"succeeded": 2, "errored": 0}, }, - request=httpx.Request( - "GET", "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" - ), + request=httpx.Request("GET", "https://api.anthropic.com/v1/messages/batches/msgbatch_abc"), ) @@ -58,9 +58,7 @@ def patched_client(): @pytest.mark.asyncio -async def test_aretrieve_batch_fires_get_with_correct_url_and_headers( - handler, patched_client -): +async def test_aretrieve_batch_fires_get_with_correct_url_and_headers(handler, patched_client): fake_client, factory = patched_client batch = await handler.aretrieve_batch( @@ -75,9 +73,7 @@ async def test_aretrieve_batch_fires_get_with_correct_url_and_headers( fake_client.get.assert_awaited_once() _, call_kwargs = fake_client.get.call_args # Exact URL built by get_retrieve_batch_url. - assert call_kwargs["url"] == ( - "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" - ) + assert call_kwargs["url"] == ("https://api.anthropic.com/v1/messages/batches/msgbatch_abc") # Auth + version + beta headers built by validate_environment. headers = call_kwargs["headers"] assert headers["x-api-key"] == "sk-ant-test" @@ -92,9 +88,7 @@ async def test_aretrieve_batch_fires_get_with_correct_url_and_headers( @pytest.mark.asyncio -async def test_aretrieve_batch_uses_anthropic_provider_for_client( - handler, patched_client -): +async def test_aretrieve_batch_uses_anthropic_provider_for_client(handler, patched_client): from litellm.types.utils import LlmProviders _, factory = patched_client @@ -110,14 +104,10 @@ async def test_aretrieve_batch_uses_anthropic_provider_for_client( @pytest.mark.asyncio -async def test_aretrieve_batch_resolves_api_key_from_model_info( - handler, patched_client -): +async def test_aretrieve_batch_resolves_api_key_from_model_info(handler, patched_client): fake_client, _ = patched_client # api_key=None -> handler falls back to AnthropicModelInfo.get_api_key(). - with patch.object( - handler.anthropic_model_info, "get_api_key", return_value="sk-from-env" - ): + with patch.object(handler.anthropic_model_info, "get_api_key", return_value="sk-from-env"): await handler.aretrieve_batch( batch_id="msgbatch_abc", api_base="https://api.anthropic.com", @@ -133,9 +123,7 @@ async def test_aretrieve_batch_resolves_api_key_from_model_info( async def test_aretrieve_batch_missing_api_key_raises(handler, patched_client): fake_client, _ = patched_client # No api_key and resolver yields None -> hard error before any network call. - with patch.object( - handler.anthropic_model_info, "get_api_key", return_value=None - ): + with patch.object(handler.anthropic_model_info, "get_api_key", return_value=None): with pytest.raises(ValueError, match="Missing Anthropic API Key"): await handler.aretrieve_batch( batch_id="msgbatch_abc", @@ -164,9 +152,7 @@ async def test_aretrieve_batch_resolves_default_api_base(handler, patched_client max_retries=0, ) _, call_kwargs = fake_client.get.call_args - assert call_kwargs["url"] == ( - "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" - ) + assert call_kwargs["url"] == ("https://api.anthropic.com/v1/messages/batches/msgbatch_abc") @pytest.mark.asyncio @@ -175,9 +161,7 @@ async def test_aretrieve_batch_raises_for_status(handler): error_response = httpx.Response( status_code=404, json={"error": "not found"}, - request=httpx.Request( - "GET", "https://api.anthropic.com/v1/messages/batches/missing" - ), + request=httpx.Request("GET", "https://api.anthropic.com/v1/messages/batches/missing"), ) fake_client = MagicMock() fake_client.get = AsyncMock(return_value=error_response) @@ -212,21 +196,15 @@ async def test_aretrieve_batch_invokes_pre_call_logging(handler, patched_client) assert pre_kwargs["input"] == "msgbatch_abc" assert pre_kwargs["api_key"] == "sk-ant-test" # The logged api_base is the full retrieve URL, not the bare base. - assert pre_kwargs["additional_args"]["api_base"] == ( - "https://api.anthropic.com/v1/messages/batches/msgbatch_abc" - ) + assert pre_kwargs["additional_args"]["api_base"] == ("https://api.anthropic.com/v1/messages/batches/msgbatch_abc") @pytest.mark.asyncio -async def test_aretrieve_batch_builds_default_logging_obj_when_absent( - handler, patched_client -): +async def test_aretrieve_batch_builds_default_logging_obj_when_absent(handler, patched_client): # logging_obj=None -> handler constructs a real Logging object; the call # must still complete (no AttributeError on a missing logger). _, _ = patched_client - with patch( - "litellm.litellm_core_utils.litellm_logging.Logging" - ) as logging_cls: + with patch("litellm.litellm_core_utils.litellm_logging.Logging") as logging_cls: logging_cls.return_value = MagicMock() batch = await handler.aretrieve_batch( batch_id="msgbatch_abc", @@ -280,3 +258,99 @@ def test_retrieve_batch_sync_runs_to_result(handler, patched_client): assert isinstance(batch, LiteLLMBatch) assert batch.id == "msgbatch_abc" assert batch.status == "completed" + + +# =========================================================================== # +# aretrieve_batch must not block the event loop on a WIF token exchange +# =========================================================================== # + +_WIF_ENV = { + "ANTHROPIC_FEDERATION_RULE_ID": "fdrl_batches_seam", + "ANTHROPIC_ORGANIZATION_ID": "org-batches-seam", + "ANTHROPIC_IDENTITY_TOKEN": "batches-seam-inline-jwt", +} + + +class _BlockingPoster: + """A token-endpoint poster that blocks until released, so the test can prove + the exchange ran off the event loop's own thread instead of freezing it.""" + + def __init__(self): + self.release = threading.Event() + self.thread_ids = [] + + def post(self, url, *, content, headers, timeout): + self.thread_ids.append(threading.get_ident()) + self.release.wait(timeout=5) + return httpx.Response( + 200, + json={ + "access_token": "sk-ant-oat01-batches-seam", + "token_type": "Bearer", + "expires_in": 3600, + }, + ) + + +@pytest.mark.asyncio +async def test_aretrieve_batch_wif_exchange_does_not_block_event_loop(handler, patched_client, monkeypatch): + """Regression: aretrieve_batch called the synchronous validate_environment + directly, so a cold WIF mint ran inline on the event loop and froze every + other concurrent coroutine until the exchange finished.""" + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + fake_client, _ = patched_client + for name in ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + ): + monkeypatch.delenv(name, raising=False) + for name, value in _WIF_ENV.items(): + monkeypatch.setenv(name, value) + + poster = _BlockingPoster() + engine = JwtBearerTokenExchangeEngine(poster=poster) + + def routed_through_injected_engine(litellm_params, api_base, model): + return get_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "get_anthropic_wif_token", routed_through_injected_engine) + + ticks = [] + + async def ticker(): + for i in range(20): + await asyncio.sleep(0.005) + ticks.append(i) + + ticker_task = asyncio.create_task(ticker()) + await asyncio.sleep(0.02) + + retrieve_task = asyncio.create_task( + handler.aretrieve_batch( + batch_id="msgbatch_abc", + api_base="https://api.anthropic.com", + api_key=None, + timeout=60.0, + max_retries=0, + ) + ) + await asyncio.sleep(0.05) + # The ticker kept advancing while the token exchange was still blocked on + # poster.release, proving the exchange did not run on the event loop. + assert len(ticks) > 0 + assert not retrieve_task.done() + + poster.release.set() + batch = await retrieve_task + await ticker_task + + assert batch.id == "msgbatch_abc" + assert poster.thread_ids + assert poster.thread_ids[0] != threading.get_ident() + sent_headers = fake_client.get.call_args.kwargs["headers"] + assert sent_headers["authorization"] == "Bearer sk-ant-oat01-batches-seam" diff --git a/tests/unit/llms/anthropic/batches/test_transformation.py b/tests/unit/llms/anthropic/batches/test_transformation.py index 419fc7740eb..26666f54241 100644 --- a/tests/unit/llms/anthropic/batches/test_transformation.py +++ b/tests/unit/llms/anthropic/batches/test_transformation.py @@ -80,8 +80,11 @@ def test_validate_environment_preserves_existing_beta_header(config): litellm_params={}, api_key="sk-ant-test", ) - # Existing beta header must NOT be overwritten. - assert headers["anthropic-beta"] == "custom-beta-value" + # Existing beta values are preserved and the batches beta is merged in. + assert set(headers["anthropic-beta"].split(",")) == { + "custom-beta-value", + "message-batches-2024-09-24", + } def test_validate_environment_oauth_key_uses_bearer(config): @@ -100,9 +103,7 @@ def test_validate_environment_oauth_key_uses_bearer(config): def test_validate_environment_missing_key_raises(config): # No api_key passed and no env credentials -> get_auth_header returns None. - with patch.object( - config.anthropic_model_info, "get_auth_header", return_value=None - ): + with patch.object(config.anthropic_model_info, "get_auth_header", return_value=None): with pytest.raises(ValueError, match="Missing Anthropic API Key"): config.validate_environment( headers={}, @@ -241,12 +242,7 @@ def test_get_retrieve_batch_url_uses_default_api_base(config): def test_transform_retrieve_batch_request_returns_empty_dict(config): - assert ( - config.transform_retrieve_batch_request( - batch_id="msgbatch_123", optional_params={}, litellm_params={} - ) - == {} - ) + assert config.transform_retrieve_batch_request(batch_id="msgbatch_123", optional_params={}, litellm_params={}) == {} # =========================================================================== # @@ -455,9 +451,7 @@ def test_transform_retrieve_response_unparseable_json_raises(config): def test_get_error_class_with_dict_headers(config): - err = config.get_error_class( - error_message="rate limited", status_code=429, headers={"x-ratelimit": "0"} - ) + err = config.get_error_class(error_message="rate limited", status_code=429, headers={"x-ratelimit": "0"}) from litellm.llms.anthropic.common_utils import AnthropicError assert isinstance(err, AnthropicError) @@ -467,9 +461,7 @@ def test_get_error_class_with_dict_headers(config): def test_get_error_class_with_httpx_headers(config): hdrs = httpx.Headers({"retry-after": "5"}) - err = config.get_error_class( - error_message="server error", status_code=500, headers=hdrs - ) + err = config.get_error_class(error_message="server error", status_code=500, headers=hdrs) assert err.status_code == 500 assert err.message == "server error" @@ -543,9 +535,7 @@ def test_transform_response_skips_malformed_lines(config): def fake_transform_parsed(*, completion_response, raw_response, model_response): mr = ModelResponse() - setattr( - mr, "usage", Usage(prompt_tokens=7, completion_tokens=3, total_tokens=10) - ) + setattr(mr, "usage", Usage(prompt_tokens=7, completion_tokens=3, total_tokens=10)) return mr with patch.object( @@ -588,13 +578,16 @@ def test_transform_response_reraises_unexpected_error(config): # A non-JSONDecodeError raised during usage aggregation must propagate # (the outer `except Exception: raise e`), not be swallowed. - with patch.object( - config.anthropic_chat_config, - "transform_parsed_response", - side_effect=fake_transform_parsed, - ), patch( - "litellm.cost_calculator.BaseTokenUsageProcessor.combine_usage_objects", - side_effect=RuntimeError("boom"), + with ( + patch.object( + config.anthropic_chat_config, + "transform_parsed_response", + side_effect=fake_transform_parsed, + ), + patch( + "litellm.cost_calculator.BaseTokenUsageProcessor.combine_usage_objects", + side_effect=RuntimeError("boom"), + ), ): with pytest.raises(RuntimeError, match="boom"): config.transform_response( diff --git a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py index 1c1b68de6d6..09921a9a85c 100644 --- a/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py +++ b/tests/unit/llms/anthropic/chat/guardrail_translation/test_anthropic_guardrail_handler.py @@ -49,9 +49,7 @@ class MockDynamicGuardrail(CustomGuardrail): input_type: Literal["request", "response"], logging_obj: Optional[Any] = None, ) -> GenericGuardrailAPIInputs: - self.dynamic_params = self.get_guardrail_dynamic_request_body_params( - request_data - ) + self.dynamic_params = self.get_guardrail_dynamic_request_body_params(request_data) return inputs @@ -198,9 +196,7 @@ class TestAnthropicMessagesHandlerStreamingRequestData: assert guardrail.request_data is not None assert guardrail.request_data["response"] is mock_response - assert ( - guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1" - ) + assert guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1" @pytest.mark.asyncio async def test_mid_stream_chunk_passes_responses_so_far_and_metadata(self): @@ -212,9 +208,7 @@ class TestAnthropicMessagesHandlerStreamingRequestData: with ( patch.object(handler, "_check_streaming_has_ended", return_value=False), - patch.object( - handler, "get_streaming_string_so_far", return_value="partial text" - ), + patch.object(handler, "get_streaming_string_so_far", return_value="partial text"), ): await handler.process_output_streaming_response( responses_so_far=responses_so_far, @@ -226,9 +220,7 @@ class TestAnthropicMessagesHandlerStreamingRequestData: assert guardrail.request_data is not None assert guardrail.request_data["responses"] is responses_so_far - assert ( - guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1" - ) + assert guardrail.request_data["litellm_metadata"]["user_api_key_user_id"] == "u-1" class TestAnthropicMessagesHandlerStreamingOutputProcessing: @@ -527,17 +519,11 @@ class TestAnthropicMessagesHandlerInputProcessing: data = { "model": "claude-3-5-sonnet-20241022", "messages": [{"role": "user", "content": "hello"}], - "litellm_metadata": { - "guardrails": [ - {"cygnal-monitor": {"extra_body": {"policy_id": "policy-123"}}} - ] - }, + "litellm_metadata": {"guardrails": [{"cygnal-monitor": {"extra_body": {"policy_id": "policy-123"}}}]}, } with patch("litellm.proxy.proxy_server.premium_user", True): - await handler.process_input_messages( - data=data, guardrail_to_apply=guardrail - ) + await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) assert data.get("litellm_metadata", {}).get("guardrails") assert guardrail.dynamic_params == {"policy_id": "policy-123"} @@ -1439,9 +1425,7 @@ class TestAnthropicMessagesHandlerInputProcessing: # Mock _check_streaming_has_ended to return False (stream not ended) with ( patch.object(handler, "_check_streaming_has_ended", return_value=False), - patch.object( - handler, "get_streaming_string_so_far", return_value="partial text" - ), + patch.object(handler, "get_streaming_string_so_far", return_value="partial text"), ): responses_so_far = [b"data: some chunk"] @@ -1472,9 +1456,7 @@ class TestAnthropicMessagesHandlerInputProcessing: data = { "model": "claude-opus-4-6", - "messages": [ - {"role": "user", "content": "What is the weather in San Francisco?"} - ], + "messages": [{"role": "user", "content": "What is the weather in San Francisco?"}], "tools": [ { "type": "tool_search_tool_regex_20251119", @@ -1637,18 +1619,14 @@ class TestAnthropicMessagesIncrementalScan: ] with patch.object(guardrail, "make_bedrock_api_request", new_callable=AsyncMock) as mock_api: mock_api.return_value = {"action": "NONE", "output": [], "outputs": []} - await handler.process_input_messages( - data=self._data(turn1, sid), guardrail_to_apply=guardrail - ) + await handler.process_input_messages(data=self._data(turn1, sid), guardrail_to_apply=guardrail) assert mock_api.call_count == 1 assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ "You are a helpful geography assistant.", "What is the capital of France?", ] mock_api.reset_mock() - await handler.process_input_messages( - data=self._data(turn2, sid), guardrail_to_apply=guardrail - ) + await handler.process_input_messages(data=self._data(turn2, sid), guardrail_to_apply=guardrail) assert mock_api.call_count == 1 assert [m["content"] for m in mock_api.call_args.kwargs["messages"]] == [ "Paris.", @@ -2087,9 +2065,7 @@ class TestAnthropicMessagesScanOnlyToolResults: await handler.process_input_messages(data=data, guardrail_to_apply=guardrail) - assert guardrail.seen_texts == ["fetched POISON page"], ( - "only the tool_result payload may reach the guardrail" - ) + assert guardrail.seen_texts == ["fetched POISON page"], "only the tool_result payload may reach the guardrail" assert guardrail.captured_inputs is not None assert guardrail.captured_inputs.get("tools") is None assert [m["role"] for m in guardrail.captured_inputs["structured_messages"]] == ["tool"] diff --git a/tests/unit/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/unit/llms/anthropic/chat/test_anthropic_chat_handler.py index 29dc6da2677..6d7d62e47aa 100644 --- a/tests/unit/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/unit/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -26,9 +26,7 @@ async def test_make_call_passes_logging_obj_to_client_post(): mock_client = AsyncMock() mock_response = MagicMock() mock_response.aiter_lines = MagicMock( - return_value=iter( - [b'data: {"type":"message_start"}\n', b'data: {"type":"message_delta"}\n'] - ) + return_value=iter([b'data: {"type":"message_start"}\n', b'data: {"type":"message_delta"}\n']) ) mock_client.post.return_value = mock_response @@ -145,9 +143,7 @@ def test_redacted_thinking_content_block_delta(): "data": "EuoBCoYBGAIiQJ/SxkPAgqxhKok29YrpJHRUJ0OT8ahCHKAwyhmRuUhtdmDX9+mn4gDzKNv3fVpQdB01zEPMzNY3QuTCd+1bdtEqQK6JuKHqdndbwpr81oVWb4wxd1GqF/7Jkw74IlQa27oobX+KuRkopr9Dllt/RDe7Se0sI1IkU7tJIAQCoP46OAwSDF51P09q67xhHlQ3ihoM2aOVlkghq/X0w8NlIjBMNvXYNbjhyrOcIg6kPFn2ed/KK7Cm5prYAtXCwkb4Wr5tUSoSHu9T5hKdJRbr6WsqEc7Lle7FULqMLZGkhqXyc3BA", }, } - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) model_response = model_response_iterator.chunk_parser(chunk=chunk) print(f"\n\nmodel_response: {model_response}\n\n") assert model_response.choices[0].delta.thinking_blocks is not None @@ -155,19 +151,14 @@ def test_redacted_thinking_content_block_delta(): print( f"\n\nmodel_response.choices[0].delta.thinking_blocks[0]: {model_response.choices[0].delta.thinking_blocks[0]}\n\n" ) - assert ( - model_response.choices[0].delta.thinking_blocks[0]["type"] - == "redacted_thinking" - ) + assert model_response.choices[0].delta.thinking_blocks[0]["type"] == "redacted_thinking" assert model_response.choices[0].delta.provider_specific_fields is not None assert "thinking_blocks" in model_response.choices[0].delta.provider_specific_fields def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -191,17 +182,12 @@ def test_streaming_thinking_blocks_are_replayable_after_signature_delta(): }, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) expected_delta_blocks = ( {"type": "thinking", "thinking": "Step 1. "}, @@ -285,9 +271,7 @@ def test_streamed_signed_thinking_round_trips_to_the_next_turn_once(): def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -307,17 +291,12 @@ def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): {"type": "content_block_stop", "index": 0}, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) assert reasoning_content == "Step 1. Step 2." @@ -328,9 +307,7 @@ def test_streaming_unsigned_thinking_deltas_keep_reasoning_content(): def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) chunks = [ { "type": "content_block_start", @@ -349,17 +326,12 @@ def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): }, ] - parsed_chunks = [ - model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks - ] + parsed_chunks = [model_response_iterator.chunk_parser(chunk=chunk) for chunk in chunks] reasoning_content = "".join( - getattr(chunk.choices[0].delta, "reasoning_content", None) or "" - for chunk in parsed_chunks + getattr(chunk.choices[0].delta, "reasoning_content", None) or "" for chunk in parsed_chunks ) thinking_blocks = tuple( - block - for chunk in parsed_chunks - for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) + block for chunk in parsed_chunks for block in (getattr(chunk.choices[0].delta, "thinking_blocks", None) or []) ) assert reasoning_content == "Step 1. Step 2." @@ -370,9 +342,7 @@ def test_streaming_truncated_thinking_deltas_keep_reasoning_content(): def test_handle_json_mode_chunk_response_format_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) response_format_tool = ChatCompletionToolCallChunk( id="tool_123", type="function", @@ -383,9 +353,7 @@ def test_handle_json_mode_chunk_response_format_tool(): index=0, ) - text, tool_use = model_response_iterator._handle_json_mode_chunk( - "", response_format_tool - ) + text, tool_use = model_response_iterator._handle_json_mode_chunk("", response_format_tool) print(f"\n\nresponse_format_tool text: {text}\n\n") print(f"\n\nresponse_format_tool tool_use: {tool_use}\n\n") @@ -394,15 +362,11 @@ def test_handle_json_mode_chunk_response_format_tool(): def test_handle_json_mode_chunk_regular_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) regular_tool = ChatCompletionToolCallChunk( id="tool_456", type="function", - function=ChatCompletionToolCallFunctionChunk( - name="get_weather", arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name="get_weather", arguments='{"location": "San Francisco, CA"}'), index=0, ) @@ -416,17 +380,13 @@ def test_handle_json_mode_chunk_regular_tool(): def test_handle_json_mode_chunk_streaming_response_format_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: response_format tool with id and name, but no arguments first_chunk = ChatCompletionToolCallChunk( id="tool_123", type="function", - function=ChatCompletionToolCallFunctionChunk( - name=RESPONSE_FORMAT_TOOL_NAME, arguments="" - ), + function=ChatCompletionToolCallFunctionChunk(name=RESPONSE_FORMAT_TOOL_NAME, arguments=""), index=0, ) @@ -434,9 +394,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): second_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments='{"question": "What is the weather?"' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments='{"question": "What is the weather?"'), index=0, ) @@ -444,9 +402,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): third_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments=', "answer": "It is sunny"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments=', "answer": "It is sunny"}'), index=0, ) @@ -477,9 +433,7 @@ def test_handle_json_mode_chunk_streaming_response_format_tool(): def test_handle_json_mode_chunk_streaming_regular_tool(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: regular tool with id and name, but no arguments first_chunk = ChatCompletionToolCallChunk( @@ -493,9 +447,7 @@ def test_handle_json_mode_chunk_streaming_regular_tool(): second_chunk = ChatCompletionToolCallChunk( id=None, type="function", - function=ChatCompletionToolCallFunctionChunk( - name=None, arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=None, arguments='{"location": "San Francisco, CA"}'), index=0, ) @@ -520,27 +472,19 @@ def test_handle_json_mode_chunk_streaming_regular_tool(): def test_response_format_tool_finish_reason(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: response_format tool response_format_tool = ChatCompletionToolCallChunk( id="tool_123", type="function", - function=ChatCompletionToolCallFunctionChunk( - name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": "test"}' - ), + function=ChatCompletionToolCallFunctionChunk(name=RESPONSE_FORMAT_TOOL_NAME, arguments='{"answer": "test"}'), index=0, ) # Process the tool call (should set converted_response_format_tool flag) - text, tool_use = model_response_iterator._handle_json_mode_chunk( - "", response_format_tool - ) - print( - f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n" - ) + text, tool_use = model_response_iterator._handle_json_mode_chunk("", response_format_tool) + print(f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n") # Simulate message_delta chunk with tool_use stop_reason message_delta_chunk = { @@ -559,25 +503,19 @@ def test_response_format_tool_finish_reason(): def test_regular_tool_finish_reason(): - model_response_iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=True - ) + model_response_iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=True) # First chunk: regular tool (not response_format) regular_tool = ChatCompletionToolCallChunk( id="tool_456", type="function", - function=ChatCompletionToolCallFunctionChunk( - name="get_weather", arguments='{"location": "San Francisco, CA"}' - ), + function=ChatCompletionToolCallFunctionChunk(name="get_weather", arguments='{"location": "San Francisco, CA"}'), index=0, ) # Process the tool call (should NOT set converted_response_format_tool flag) text, tool_use = model_response_iterator._handle_json_mode_chunk("", regular_tool) - print( - f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n" - ) + print(f"\n\nconverted_response_format_tool flag: {model_response_iterator.converted_response_format_tool}\n\n") # Simulate message_delta chunk with tool_use stop_reason message_delta_chunk = { @@ -637,9 +575,7 @@ def test_text_only_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, f"Expected index=0, got {parsed.choices[0].index}" def test_message_delta_without_usage_returns_chunk_with_no_usage(): @@ -830,9 +766,7 @@ def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinkin ] self._write_response( content_type="text/event-stream", - body="".join( - f"data: {json.dumps(event)}\n\n" for event in events - ).encode("utf-8"), + body="".join(f"data: {json.dumps(event)}\n\n" for event in events).encode("utf-8"), ) return @@ -913,13 +847,9 @@ def test_anthropic_completion_streaming_usage_matches_non_streaming_with_thinkin assert content_chunks == [answer_text] assert stream_usage is not None stream_completion_details = stream_usage["completion_tokens_details"] - assert ( - stream_completion_details["reasoning_tokens"] - == non_stream_details.reasoning_tokens - ) + assert stream_completion_details["reasoning_tokens"] == non_stream_details.reasoning_tokens assert stream_completion_details["text_tokens"] == ( - stream_usage["completion_tokens"] - - stream_completion_details["reasoning_tokens"] + stream_usage["completion_tokens"] - stream_completion_details["reasoning_tokens"] ) assert requests_seen == [ { @@ -1011,9 +941,9 @@ def test_text_and_tool_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0 for chunk type {chunk.get('type')}, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, ( + f"Expected index=0 for chunk type {chunk.get('type')}, got {parsed.choices[0].index}" + ) def test_multiple_tools_streaming_has_index_zero(): @@ -1066,15 +996,11 @@ def test_multiple_tools_streaming_has_index_zero(): for chunk in chunks: parsed = iterator.chunk_parser(chunk) if parsed.choices: - assert ( - parsed.choices[0].index == 0 - ), f"Expected index=0, got {parsed.choices[0].index}" + assert parsed.choices[0].index == 0, f"Expected index=0, got {parsed.choices[0].index}" def test_streaming_chunks_have_stable_ids(): - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) first_chunk = { "type": "content_block_delta", "index": 0, @@ -1099,9 +1025,7 @@ def test_partial_json_chunk_accumulation(): This tests the fix for https://github.com/BerriAI/litellm/issues/17473 where network fragmentation can cause SSE data to arrive in partial chunks. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) partial_chunk_1 = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hel' partial_chunk_2 = 'lo"}}' @@ -1109,31 +1033,21 @@ def test_partial_json_chunk_accumulation(): # First partial chunk should return None (still accumulating) result1 = iterator._parse_sse_data(f"data:{partial_chunk_1}") assert result1 is None, "First partial chunk should return None while accumulating" - assert ( - iterator.chunk_type == "accumulated_json" - ), "Should switch to accumulated_json mode" - assert ( - iterator.accumulated_json == partial_chunk_1 - ), "Should have accumulated first part" + assert iterator.chunk_type == "accumulated_json", "Should switch to accumulated_json mode" + assert iterator.accumulated_json == partial_chunk_1, "Should have accumulated first part" # Second partial chunk should complete the JSON and return a parsed result result2 = iterator._parse_sse_data(f"data:{partial_chunk_2}") assert result2 is not None, "Second chunk should return parsed result" - assert ( - iterator.accumulated_json == "" - ), "Buffer should be cleared after successful parse" - assert ( - result2.choices[0].delta.content == "Hello" - ), f"Expected 'Hello', got '{result2.choices[0].delta.content}'" + assert iterator.accumulated_json == "", "Buffer should be cleared after successful parse" + assert result2.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result2.choices[0].delta.content}'" def test_complete_json_chunk_no_accumulation(): """ Test that complete JSON chunks are parsed immediately without accumulation. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) complete_chunk = '{"type":"content_block_delta","index":0,"delta":{"type":"text_delta","text":"Hello"}}' @@ -1141,18 +1055,14 @@ def test_complete_json_chunk_no_accumulation(): assert result is not None, "Complete chunk should return parsed result immediately" assert iterator.chunk_type == "valid_json", "Should remain in valid_json mode" assert iterator.accumulated_json == "", "Buffer should remain empty" - assert ( - result.choices[0].delta.content == "Hello" - ), f"Expected 'Hello', got '{result.choices[0].delta.content}'" + assert result.choices[0].delta.content == "Hello", f"Expected 'Hello', got '{result.choices[0].delta.content}'" def test_multiple_partial_chunks_accumulation(): """ Test that multiple partial chunks can be accumulated across several iterations. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Split a JSON chunk into three parts part1 = '{"type":"content_block_del' @@ -1320,9 +1230,7 @@ def test_web_search_tool_result_no_extra_tool_calls(): The issue was that web_search_tool_result blocks have input_json_delta events with {} that were incorrectly being converted to tool calls. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence: # 1. server_tool_use block starts (web_search) @@ -1397,9 +1305,7 @@ def test_web_search_tool_result_no_extra_tool_calls(): # Should have exactly 2 tool calls: # 1. From content_block_start (server_tool_use) with id and name # 2. From content_block_delta with the actual query - assert ( - len(tool_calls_emitted) == 2 - ), f"Expected 2 tool calls, got {len(tool_calls_emitted)}" + assert len(tool_calls_emitted) == 2, f"Expected 2 tool calls, got {len(tool_calls_emitted)}" # First tool call should have the id and name assert tool_calls_emitted[0]["id"] == "srvtoolu_01ABC123" @@ -1415,9 +1321,7 @@ def test_current_content_block_type_tracking(): """ Test that current_content_block_type is properly tracked and reset. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Initially should be None assert iterator.current_content_block_type is None @@ -1516,9 +1420,7 @@ def test_web_search_tool_result_captured_in_provider_specific_fields(): The web_search_tool_result content comes ALL AT ONCE in content_block_start, not in deltas, so we need to capture it there. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence with web_search_tool_result chunks = [ @@ -1589,23 +1491,15 @@ def test_web_search_tool_result_captured_in_provider_specific_fields(): and parsed.choices[0].delta.provider_specific_fields and "web_search_results" in parsed.choices[0].delta.provider_specific_fields ): - web_search_results = parsed.choices[0].delta.provider_specific_fields[ - "web_search_results" - ] + web_search_results = parsed.choices[0].delta.provider_specific_fields["web_search_results"] # Verify web_search_results was captured assert web_search_results is not None, "web_search_results should be captured" assert len(web_search_results) == 1, "Should have 1 web_search_tool_result block" - assert ( - web_search_results[0]["type"] == "web_search_tool_result" - ), "Block type should be web_search_tool_result" - assert ( - web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123" - ), "tool_use_id should match" + assert web_search_results[0]["type"] == "web_search_tool_result", "Block type should be web_search_tool_result" + assert web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123", "tool_use_id should match" assert len(web_search_results[0]["content"]) == 2, "Should have 2 search results" - assert ( - web_search_results[0]["content"][0]["title"] == "Fun Otter Facts" - ), "First result title should match" + assert web_search_results[0]["content"][0]["title"] == "Fun Otter Facts", "First result title should match" def test_web_fetch_tool_result_captured_in_provider_specific_fields(): @@ -1619,9 +1513,7 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields(): The web_fetch_tool_result content comes ALL AT ONCE in content_block_start, not in deltas, so we need to capture it there. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate the streaming sequence with web_fetch_tool_result chunks = [ @@ -1692,25 +1584,15 @@ def test_web_fetch_tool_result_captured_in_provider_specific_fields(): and parsed.choices[0].delta.provider_specific_fields and "web_search_results" in parsed.choices[0].delta.provider_specific_fields ): - web_search_results = parsed.choices[0].delta.provider_specific_fields[ - "web_search_results" - ] + web_search_results = parsed.choices[0].delta.provider_specific_fields["web_search_results"] # Verify web_fetch_tool_result was captured (stored in web_search_results list) assert web_search_results is not None, "web_search_results should be captured" assert len(web_search_results) == 1, "Should have 1 web_fetch_tool_result block" - assert ( - web_search_results[0]["type"] == "web_fetch_tool_result" - ), "Block type should be web_fetch_tool_result" - assert ( - web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123" - ), "tool_use_id should match" - assert ( - web_search_results[0]["content"]["url"] == "https://example.com" - ), "URL should match" - assert ( - web_search_results[0]["content"]["content"]["title"] == "Example Page" - ), "Title should match" + assert web_search_results[0]["type"] == "web_fetch_tool_result", "Block type should be web_fetch_tool_result" + assert web_search_results[0]["tool_use_id"] == "srvtoolu_01ABC123", "tool_use_id should match" + assert web_search_results[0]["content"]["url"] == "https://example.com", "URL should match" + assert web_search_results[0]["content"]["content"]["title"] == "Example Page", "Title should match" def test_web_fetch_tool_result_no_extra_tool_calls(): @@ -1723,9 +1605,7 @@ def test_web_fetch_tool_result_no_extra_tool_calls(): The issue was that web_fetch_tool_result blocks have input_json_delta events with {} that were incorrectly being converted to tool calls. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # to verify it doesn't emit tool calls chunks = [ @@ -1769,9 +1649,9 @@ def test_web_fetch_tool_result_no_extra_tool_calls(): tool_call_count += 1 # Should have 0 tool calls - web_fetch_tool_result should not emit tool calls - assert ( - tool_call_count == 0 - ), f"Expected 0 tool calls, got {tool_call_count}. web_fetch_tool_result should not emit tool calls" + assert tool_call_count == 0, ( + f"Expected 0 tool calls, got {tool_call_count}. web_fetch_tool_result should not emit tool calls" + ) def test_container_in_provider_specific_fields_streaming(): @@ -1781,9 +1661,7 @@ def test_container_in_provider_specific_fields_streaming(): When container with skills is used, the container field should be present in the provider_specific_fields of the message_delta chunk. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=True, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=True, json_mode=False) # Simulate streaming chunks chunks = [ @@ -1851,20 +1729,12 @@ def test_container_in_provider_specific_fields_streaming(): and parsed.choices[0].delta.provider_specific_fields and "container" in parsed.choices[0].delta.provider_specific_fields ): - container_field = parsed.choices[0].delta.provider_specific_fields[ - "container" - ] + container_field = parsed.choices[0].delta.provider_specific_fields["container"] # Verify container was captured - assert ( - container_field is not None - ), "container should be captured in provider_specific_fields" - assert ( - container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p" - ), "container id should match" - assert ( - container_field["expires_at"] == "2025-12-16T04:57:16.913181Z" - ), "expires_at should match" + assert container_field is not None, "container should be captured in provider_specific_fields" + assert container_field["id"] == "container_011CW9hA9zpZ8xD3bjjShy4p", "container id should match" + assert container_field["expires_at"] == "2025-12-16T04:57:16.913181Z", "expires_at should match" assert len(container_field["skills"]) == 1, "Should have 1 skill" assert container_field["skills"][0]["skill_id"] == "pptx", "skill_id should be pptx" assert container_field["skills"][0]["version"] == "20251013", "version should match" @@ -1877,9 +1747,7 @@ def test_container_in_provider_specific_fields_non_streaming(): When container with skills is used in non-streaming, the container field should be present in the provider_specific_fields of the response. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) # Simulate a message_delta chunk with container (as it would appear in non-streaming) message_delta_chunk = { @@ -1915,21 +1783,13 @@ def test_container_in_provider_specific_fields_non_streaming(): # Verify container is in provider_specific_fields assert model_response.choices[0].delta.provider_specific_fields is not None assert "container" in model_response.choices[0].delta.provider_specific_fields - container_field = model_response.choices[0].delta.provider_specific_fields[ - "container" - ] + container_field = model_response.choices[0].delta.provider_specific_fields["container"] assert container_field["id"] == "container_abc123xyz", "container id should match" - assert ( - container_field["expires_at"] == "2025-12-20T10:30:00.000000Z" - ), "expires_at should match" + assert container_field["expires_at"] == "2025-12-20T10:30:00.000000Z", "expires_at should match" assert len(container_field["skills"]) == 2, "Should have 2 skills" - assert ( - container_field["skills"][0]["skill_id"] == "code_execution" - ), "First skill_id should be code_execution" - assert ( - container_field["skills"][1]["skill_id"] == "pptx" - ), "Second skill_id should be pptx" + assert container_field["skills"][0]["skill_id"] == "code_execution", "First skill_id should be code_execution" + assert container_field["skills"][1]["skill_id"] == "pptx", "Second skill_id should be pptx" def test_container_absent_when_not_provided(): @@ -1938,9 +1798,7 @@ def test_container_absent_when_not_provided(): This ensures we don't add empty or None container fields. """ - iterator = ModelResponseIterator( - streaming_response=MagicMock(), sync_stream=False, json_mode=False - ) + iterator = ModelResponseIterator(streaming_response=MagicMock(), sync_stream=False, json_mode=False) # message_delta without container message_delta_chunk = { @@ -1959,9 +1817,9 @@ def test_container_absent_when_not_provided(): # Verify container is NOT in provider_specific_fields when not provided if model_response.choices[0].delta.provider_specific_fields: - assert ( - "container" not in model_response.choices[0].delta.provider_specific_fields - ), "container should not be present when not provided in delta" + assert "container" not in model_response.choices[0].delta.provider_specific_fields, ( + "container should not be present when not provided in delta" + ) def test_streaming_code_execution_produces_code_interpreter_results(): @@ -2157,8 +2015,7 @@ def test_streaming_multiple_code_executions_no_duplicates(): # Second (final) emission: cumulative list with BOTH results # This is what stream_chunk_builder will pick as "last value wins" assert len(emissions[1]) == 2, ( - f"Expected final emission to have 2 results, got {len(emissions[1])}. " - f"IDs: {[r.id for r in emissions[1]]}" + f"Expected final emission to have 2 results, got {len(emissions[1])}. IDs: {[r.id for r in emissions[1]]}" ) assert emissions[1][0].id == "srvtoolu_01AAA" assert emissions[1][0].code == "echo first" @@ -2322,9 +2179,7 @@ def test_empty_output_produces_null_outputs(): assert code_results is not None, "No code_interpreter_results emitted" assert len(code_results) == 1 assert code_results[0].id == "srvtoolu_01AAA" - assert ( - code_results[0].outputs is None - ), f"Expected outputs=None for empty execution, got {code_results[0].outputs}" + assert code_results[0].outputs is None, f"Expected outputs=None for empty execution, got {code_results[0].outputs}" def test_non_bash_tool_result_skipped(): @@ -2387,12 +2242,10 @@ def test_non_bash_tool_result_skipped(): code_results = psf["code_interpreter_results"] # code_interpreter_results should be emitted but empty (no bash results) - assert ( - code_results is not None - ), "Expected code_interpreter_results key to be emitted" - assert ( - len(code_results) == 0 - ), f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" + assert code_results is not None, "Expected code_interpreter_results key to be emitted" + assert len(code_results) == 0, ( + f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" + ) class TestAnthropicChatCompletionPreCallLogging: diff --git a/tests/unit/llms/anthropic/chat/test_code_interpreter_results_extraction.py b/tests/unit/llms/anthropic/chat/test_code_interpreter_results_extraction.py index 60e45c9b8ce..99266f92c56 100644 --- a/tests/unit/llms/anthropic/chat/test_code_interpreter_results_extraction.py +++ b/tests/unit/llms/anthropic/chat/test_code_interpreter_results_extraction.py @@ -136,22 +136,14 @@ def test_in_place_substitution_preserves_ordering(): responses_output = [msg_item, fc_exec1, fc_regular, fc_exec2] # Apply the same logic as _transform_chat_completion_choices_to_responses_output - tool_result_items = ( - LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp) - ) + tool_result_items = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(resp) if tool_result_items: - result_by_id = { - (item.get("id") if isinstance(item, dict) else item.id): item - for item in tool_result_items - } + result_by_id = {(item.get("id") if isinstance(item, dict) else item.id): item for item in tool_result_items} replaced_ids = set(result_by_id.keys()) responses_output = [ ( result_by_id[getattr(item, "call_id", None)] - if ( - getattr(item, "type", None) == "function_call" - and getattr(item, "call_id", None) in replaced_ids - ) + if (getattr(item, "type", None) == "function_call" and getattr(item, "call_id", None) in replaced_ids) else item ) for item in responses_output @@ -255,9 +247,7 @@ def test_end_to_end_streaming_chunks_to_code_interpreter_output(): assert code_results[0]["code"] == "echo e2e_test" # Step 3: Extract via _extract_tool_result_output_items (Responses API layer) - tool_result_items = ( - LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(assembled) - ) + tool_result_items = LiteLLMCompletionResponsesConfig._extract_tool_result_output_items(assembled) assert len(tool_result_items) == 1 item = tool_result_items[0] # Items are reconstructed as Pydantic OutputCodeInterpreterCall objects diff --git a/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py b/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py index d9763f173a7..f385c2f2211 100644 --- a/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py +++ b/tests/unit/llms/anthropic/files/test_anthropic_files_transformation.py @@ -5,7 +5,9 @@ Tests the AnthropicFilesConfig class which transforms between OpenAI-compatible file operations and Anthropic's Files API format. """ +import asyncio import io +import threading import time import httpx @@ -90,6 +92,38 @@ class TestAnthropicFilesConfig: api_key=None, ) + @pytest.mark.asyncio + async def test_avalidate_environment_sets_headers(self): + headers = {} + result = await self.config.avalidate_environment( + headers=headers, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key="sk-ant-test-key", + ) + assert result["x-api-key"] == "sk-ant-test-key" + assert result["anthropic-version"] == "2023-06-01" + assert result["anthropic-beta"] == ANTHROPIC_FILES_BETA_HEADER + + @pytest.mark.asyncio + @patch.dict("os.environ", {}, clear=True) + @patch( + "litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key", + return_value=None, + ) + async def test_avalidate_environment_missing_api_key(self, mock_get_key): + with pytest.raises(ValueError, match="Anthropic API key is required"): + await self.config.avalidate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + def test_get_supported_openai_params(self): params = self.config.get_supported_openai_params(model="") assert "purpose" in params @@ -187,10 +221,7 @@ class TestAnthropicFilesConfig: litellm_params={}, ) - assert ( - url - == f"{ANTHROPIC_FILES_API_BASE}/v1/files/..%2F..%2Fv1%2Fmessages%2Fbatches%3Flimit%3D1%23frag" - ) + assert url == f"{ANTHROPIC_FILES_API_BASE}/v1/files/..%2F..%2Fv1%2Fmessages%2Fbatches%3Flimit%3D1%23frag" assert params == {} def test_transform_retrieve_file_response(self): @@ -411,6 +442,108 @@ class TestAnthropicFilesConfig: assert error.message == "Not found" +_WIF_ENV = { + "ANTHROPIC_FEDERATION_RULE_ID": "fdrl_files_seam", + "ANTHROPIC_ORGANIZATION_ID": "org-files-seam", + "ANTHROPIC_IDENTITY_TOKEN": "files-seam-inline-jwt", +} + + +class _BlockingPoster: + """A token-endpoint poster that blocks until released, so the test can prove + the exchange ran off the event loop's own thread instead of freezing it.""" + + def __init__(self): + self.release = threading.Event() + self.thread_ids = [] + + def post(self, url, *, content, headers, timeout): + self.thread_ids.append(threading.get_ident()) + self.release.wait(timeout=5) + return httpx.Response( + 200, + json={ + "access_token": "sk-ant-oat01-files-seam", + "token_type": "Bearer", + "expires_in": 3600, + }, + ) + + +class TestAnthropicFilesConfigWifAsyncSeam: + """Regression (Greptile P1): avalidate_environment must resolve workload identity + federation through the async token-exchange facade, never the blocking sync one, + so a cold WIF mint on async file retrieval doesn't freeze the event loop.""" + + def setup_method(self): + self.config = AnthropicFilesConfig() + + @pytest.mark.asyncio + async def test_avalidate_environment_wif_exchange_does_not_block_event_loop(self, monkeypatch): + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import aget_anthropic_wif_token, get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + for name in ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + ): + monkeypatch.delenv(name, raising=False) + for name, value in _WIF_ENV.items(): + monkeypatch.setenv(name, value) + + poster = _BlockingPoster() + engine = JwtBearerTokenExchangeEngine(poster=poster) + sync_calls = [] + + def sync_shim(litellm_params, api_base, model): + sync_calls.append(model) + return get_anthropic_wif_token(litellm_params, api_base, model, engine) + + async def async_shim(litellm_params, api_base, model): + return await aget_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "get_anthropic_wif_token", sync_shim) + monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim) + + ticks = [] + + async def ticker(): + for i in range(20): + await asyncio.sleep(0.005) + ticks.append(i) + + ticker_task = asyncio.create_task(ticker()) + await asyncio.sleep(0.02) + + validate_task = asyncio.create_task( + self.config.avalidate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + ) + ) + await asyncio.sleep(0.05) + # The ticker kept advancing while the exchange was still blocked on + # poster.release, proving avalidate_environment did not run it inline. + assert len(ticks) > 0 + assert not validate_task.done() + + poster.release.set() + headers = await validate_task + await ticker_task + + assert headers["authorization"] == "Bearer sk-ant-oat01-files-seam" + assert sync_calls == [] + assert poster.thread_ids + assert poster.thread_ids[0] != threading.get_ident() + + class TestProviderConfigRegistration: """Test that AnthropicFilesConfig is properly registered.""" diff --git a/tests/unit/llms/anthropic/messages/test_advisor_orchestration.py b/tests/unit/llms/anthropic/messages/test_advisor_orchestration.py index fc80a285ec6..91647472f81 100644 --- a/tests/unit/llms/anthropic/messages/test_advisor_orchestration.py +++ b/tests/unit/llms/anthropic/messages/test_advisor_orchestration.py @@ -101,9 +101,7 @@ async def test_anthropic_native_interceptor_skipped(): ) h = AdvisorOrchestrationHandler() - assert not h.can_handle( - [ADVISOR_TOOL], "anthropic" - ), "Interceptor must NOT trigger for anthropic provider" + assert not h.can_handle([ADVISOR_TOOL], "anthropic"), "Interceptor must NOT trigger for anthropic provider" # --------------------------------------------------------------------------- @@ -204,9 +202,7 @@ async def test_loop_one_advisor_call(): assert "is_prime" in texts[0]["text"] # No advisor tool_use blocks in final response - advisor_uses = [ - b for b in content if b.get("type") == "tool_use" and b.get("name") == "advisor" - ] + advisor_uses = [b for b in content if b.get("type") == "tool_use" and b.get("name") == "advisor"] assert len(advisor_uses) == 0 @@ -366,9 +362,7 @@ async def test_prior_advisor_blocks_replaced_in_history(): # Text block with advisor feedback must be present text_blocks = [b for b in content if b.get("type") == "text"] - feedback_blocks = [ - b for b in text_blocks if "advisor_feedback" in b.get("text", "") - ] + feedback_blocks = [b for b in text_blocks if "advisor_feedback" in b.get("text", "")] assert len(feedback_blocks) >= 1 assert "trial division" in feedback_blocks[0]["text"] @@ -707,11 +701,7 @@ async def test_advisor_ignores_tool_credentials_when_clientside_disabled(): with patch.dict( sys.modules, - { - "litellm.proxy.proxy_server": _fake_proxy_server( - {"allow_client_side_credentials": False} - ) - }, + {"litellm.proxy.proxy_server": _fake_proxy_server({"allow_client_side_credentials": False})}, ): captured = await _run_advisor_and_capture_subcall_kwargs() assert captured["api_key"] is None @@ -726,11 +716,7 @@ async def test_advisor_uses_tool_credentials_when_clientside_enabled(): with patch.dict( sys.modules, - { - "litellm.proxy.proxy_server": _fake_proxy_server( - {"allow_client_side_credentials": True} - ) - }, + {"litellm.proxy.proxy_server": _fake_proxy_server({"allow_client_side_credentials": True})}, ): captured = await _run_advisor_and_capture_subcall_kwargs() assert captured["api_key"] == "sk-other" diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_handler_output_config_passthrough.py b/tests/unit/llms/anthropic/pass_through/adapters/test_handler_output_config_passthrough.py index 2dc1202a8c8..282373d4b6d 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_handler_output_config_passthrough.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_handler_output_config_passthrough.py @@ -34,9 +34,7 @@ import pytest # Anchor sys.path to this file's location — not the working-directory-relative # pattern Greptile flagged on PR #23706. Resolves correctly regardless of # where pytest is invoked from. -sys.path.insert( - 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../../..")) -) +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../../.."))) from litellm.llms.anthropic.pass_through.adapters.handler import ( ANTHROPIC_ONLY_REQUEST_KEYS, @@ -187,9 +185,7 @@ class TestOutputConfigStrippedFromCompletionKwargs: result = _call_prepare( extra_kwargs={ "custom_llm_provider": "azure", - "output_config": { - "format": {"type": "json_schema", "schema": losing_schema} - }, + "output_config": {"format": {"type": "json_schema", "schema": losing_schema}}, }, output_format={"type": "json_schema", "schema": winning_schema}, ) diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_combined_chunk.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_combined_chunk.py index 5a7cf652b95..9a6940ff12b 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_combined_chunk.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_combined_chunk.py @@ -31,9 +31,7 @@ from litellm.types.utils import ( ) -def _build_fake_stream( - content: str, finish_reason: str = "stop" -) -> MockResponseIterator: +def _build_fake_stream(content: str, finish_reason: str = "stop") -> MockResponseIterator: """Mimic a Vertex Gemma `:predict` fake stream: one collapsed chunk.""" model_response = ModelResponse() model_response.choices = [ @@ -133,9 +131,7 @@ def test_delayed_usage_chunk_preserves_cache_tokens(): wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="gpt-4o") events = list(wrapper) - message_delta = next( - event for event in events if event.get("type") == "message_delta" - ) + message_delta = next(event for event in events if event.get("type") == "message_delta") assert message_delta["usage"]["input_tokens"] == 70 assert message_delta["usage"]["output_tokens"] == 5 @@ -145,13 +141,7 @@ def test_delayed_usage_chunk_preserves_cache_tokens(): def test_splitter_passes_through_non_combined_chunks(): """A chunk with content but no finish_reason is not split.""" - chunk = ModelResponseStream( - choices=[ - StreamingChoices( - index=0, delta=Delta(content="partial"), finish_reason=None - ) - ] - ) + chunk = ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content="partial"), finish_reason=None)]) chunks = list(_CombinedChunkSplitter(iter([chunk]))) assert len(chunks) == 1 assert chunks[0].choices[0].delta.content == "partial" @@ -159,11 +149,7 @@ def test_splitter_passes_through_non_combined_chunks(): def test_splitter_splits_combined_chunk_into_content_then_finish(): """A chunk with both content and finish_reason becomes two chunks.""" - chunk = ModelResponseStream( - choices=[ - StreamingChoices(index=0, delta=Delta(content="done"), finish_reason="stop") - ] - ) + chunk = ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content="done"), finish_reason="stop")]) content_chunk, finish_chunk = list(_CombinedChunkSplitter(iter([chunk]))) assert content_chunk.choices[0].delta.content == "done" @@ -193,9 +179,7 @@ def test_split_clears_reasoning_and_thinking_on_finish_chunk(): reasoning_content="some reasoning", thinking_blocks=[{"type": "thinking"}], ) - chunk = SimpleNamespace( - choices=[SimpleNamespace(finish_reason="stop", delta=delta)] - ) + chunk = SimpleNamespace(choices=[SimpleNamespace(finish_reason="stop", delta=delta)]) content_chunk, finish_chunk = _CombinedChunkSplitter._split(chunk) diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_compaction.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_compaction.py index 3f6587b9338..2610824edf4 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_compaction.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_compaction.py @@ -22,9 +22,7 @@ def _make_text_chunk( StreamingChoices( finish_reason=finish_reason, index=0, - delta=Delta( - content=text, role="assistant" if text else None, tool_calls=None - ), + delta=Delta(content=text, role="assistant" if text else None, tool_calls=None), logprobs=None, ) ] @@ -73,34 +71,23 @@ async def test_stream_emits_compaction_block_before_text(): compaction_start = next( e for e in events - if e.get("type") == "content_block_start" - and e.get("content_block", {}).get("type") == "compaction" + if e.get("type") == "content_block_start" and e.get("content_block", {}).get("type") == "compaction" ) assert compaction_start["index"] == 0 compaction_delta = next( e for e in events - if e.get("type") == "content_block_delta" - and e.get("delta", {}).get("type") == "compaction_delta" + if e.get("type") == "content_block_delta" and e.get("delta", {}).get("type") == "compaction_delta" ) assert compaction_delta["index"] == 0 - assert ( - compaction_delta["delta"]["content"] == "Summary of prior conversation turns." - ) + assert compaction_delta["delta"]["content"] == "Summary of prior conversation turns." - compaction_stop = next( - e - for e in events - if e.get("type") == "content_block_stop" and e.get("index") == 0 - ) + compaction_stop = next(e for e in events if e.get("type") == "content_block_stop" and e.get("index") == 0) assert compaction_stop is not None text_start = next( - e - for e in events - if e.get("type") == "content_block_start" - and e.get("content_block", {}).get("type") == "text" + e for e in events if e.get("type") == "content_block_start" and e.get("content_block", {}).get("type") == "text" ) assert text_start["index"] == 1 @@ -177,14 +164,9 @@ async def test_stream_without_compaction_block_unchanged(): events = await _collect_events_async(wrapper) assert not any( - e.get("content_block", {}).get("type") == "compaction" - for e in events - if e.get("type") == "content_block_start" + e.get("content_block", {}).get("type") == "compaction" for e in events if e.get("type") == "content_block_start" ) text_start = next( - e - for e in events - if e.get("type") == "content_block_start" - and e.get("content_block", {}).get("type") == "text" + e for e in events if e.get("type") == "content_block_start" and e.get("content_block", {}).get("type") == "text" ) assert text_start["index"] == 0 diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_empty_choices.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_empty_choices.py index ca2532fce56..66825124f25 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_empty_choices.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_empty_choices.py @@ -19,9 +19,7 @@ from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices, Us def _text_chunk(text: str) -> ModelResponseStream: - return ModelResponseStream( - choices=[StreamingChoices(index=0, delta=Delta(content=text), finish_reason=None)] - ) + return ModelResponseStream(choices=[StreamingChoices(index=0, delta=Delta(content=text), finish_reason=None)]) def _finish_chunk() -> ModelResponseStream: @@ -61,9 +59,7 @@ def test_leading_metadata_chunk_without_choices_does_not_kill_stream(): wrapper = AnthropicStreamWrapper(completion_stream=iter(chunks), model="mock-model") events = list(wrapper) - text = "".join( - event["delta"]["text"] for event in events if event.get("type") == "content_block_delta" - ) + text = "".join(event["delta"]["text"] for event in events if event.get("type") == "content_block_delta") assert text == "Hello there" assert events[-1]["type"] == "message_stop" diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_first_delta.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_first_delta.py index 18cf42776f9..92234282bb2 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_first_delta.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_first_delta.py @@ -637,11 +637,7 @@ def _thinking_first_chunks() -> List[MagicMock]: def _assert_thinking_first_block_opens_at_index_zero(events: List[dict]) -> None: - starts = [ - (e["index"], e["content_block"]["type"]) - for e in events - if e.get("type") == "content_block_start" - ] + starts = [(e["index"], e["content_block"]["type"]) for e in events if e.get("type") == "content_block_start"] assert starts == [(0, "thinking"), (1, "text")], starts assert "" not in _text_deltas(events) assert _thinking_deltas(events) == ["Let me think", "about it."] @@ -1110,9 +1106,7 @@ def test_tool_block_start_emitted_without_awaiting_the_next_chunk_sync(): "name": "Write", "input": {}, } - assert stream.pulled == 1, ( - f"content_block_start was withheld until {stream.pulled} upstream chunks had arrived" - ) + assert stream.pulled == 1, f"content_block_start was withheld until {stream.pulled} upstream chunks had arrived" @pytest.mark.asyncio @@ -1127,9 +1121,7 @@ async def test_tool_block_start_emitted_without_awaiting_the_next_chunk_async(): start = await wrapper.__anext__() assert start["type"] == "content_block_start" assert start["content_block"]["name"] == "Write" - assert stream.pulled == 1, ( - f"content_block_start was withheld until {stream.pulled} upstream chunks had arrived" - ) + assert stream.pulled == 1, f"content_block_start was withheld until {stream.pulled} upstream chunks had arrived" @pytest.mark.parametrize("is_async", [False, True]) diff --git a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_tool_args.py b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_tool_args.py index e9fe65ec8b0..d90472601ff 100644 --- a/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_tool_args.py +++ b/tests/unit/llms/anthropic/pass_through/adapters/test_streaming_iterator_tool_args.py @@ -129,15 +129,13 @@ async def test_async_stream_emits_input_json_delta_for_bundled_tool_args(): ): input_json_delta_idx = i - assert ( - tool_start_idx is not None - ), f"Expected content_block_start with type=tool_use; events: {event_types}" - assert ( - input_json_delta_idx is not None - ), f"Expected content_block_delta with input_json_delta; events: {event_types}" - assert ( - input_json_delta_idx == tool_start_idx + 1 - ), "input_json_delta should immediately follow the tool_use content_block_start" + assert tool_start_idx is not None, f"Expected content_block_start with type=tool_use; events: {event_types}" + assert input_json_delta_idx is not None, ( + f"Expected content_block_delta with input_json_delta; events: {event_types}" + ) + assert input_json_delta_idx == tool_start_idx + 1, ( + "input_json_delta should immediately follow the tool_use content_block_start" + ) # Verify the delta carries the tool arguments delta_event = events[input_json_delta_idx] @@ -230,8 +228,7 @@ async def test_async_stream_no_extra_delta_when_tool_args_empty(): and e["delta"].get("type") == "input_json_delta" ] assert len(input_json_deltas) == 1, ( - f"Expected exactly 1 input_json_delta (from the follow-up chunk), " - f"got {len(input_json_deltas)}" + f"Expected exactly 1 input_json_delta (from the follow-up chunk), got {len(input_json_deltas)}" ) assert input_json_deltas[0]["delta"]["partial_json"] == '{"location": "NYC"}' @@ -291,15 +288,13 @@ def test_sync_stream_emits_input_json_delta_for_bundled_tool_args(): ): input_json_delta_idx = i - assert ( - tool_start_idx is not None - ), f"Expected content_block_start with type=tool_use; events: {event_types}" - assert ( - input_json_delta_idx is not None - ), f"Expected content_block_delta with input_json_delta; events: {event_types}" - assert ( - input_json_delta_idx == tool_start_idx + 1 - ), "input_json_delta should immediately follow the tool_use content_block_start" + assert tool_start_idx is not None, f"Expected content_block_start with type=tool_use; events: {event_types}" + assert input_json_delta_idx is not None, ( + f"Expected content_block_delta with input_json_delta; events: {event_types}" + ) + assert input_json_delta_idx == tool_start_idx + 1, ( + "input_json_delta should immediately follow the tool_use content_block_start" + ) assert json.loads(events[input_json_delta_idx]["delta"]["partial_json"]) == {"location": "Boston"} @@ -343,9 +338,7 @@ def test_sync_stream_no_extra_delta_when_tool_args_empty(): ) wrapper = AnthropicStreamWrapper( - completion_stream=iter( - [text_chunk, tool_name_chunk, tool_args_chunk, finish_chunk] - ), + completion_stream=iter([text_chunk, tool_name_chunk, tool_args_chunk, finish_chunk]), model="test-model", ) @@ -374,7 +367,6 @@ def test_sync_stream_no_extra_delta_when_tool_args_empty(): and e["delta"].get("type") == "input_json_delta" ] assert len(input_json_deltas) == 1, ( - f"Expected exactly 1 input_json_delta (from the follow-up chunk), " - f"got {len(input_json_deltas)}" + f"Expected exactly 1 input_json_delta (from the follow-up chunk), got {len(input_json_deltas)}" ) assert input_json_deltas[0]["delta"]["partial_json"] == '{"location": "NYC"}' diff --git a/tests/unit/llms/anthropic/pass_through/context_management/test_clear_tool_uses.py b/tests/unit/llms/anthropic/pass_through/context_management/test_clear_tool_uses.py index 7a4a0f40ecc..ee50bdff461 100644 --- a/tests/unit/llms/anthropic/pass_through/context_management/test_clear_tool_uses.py +++ b/tests/unit/llms/anthropic/pass_through/context_management/test_clear_tool_uses.py @@ -258,9 +258,7 @@ def test_tool_result_list_content_shape_preserved(): {"role": "user", "content": "Hi"}, { "role": "assistant", - "content": [ - {"type": "tool_use", "id": "toolu_a", "name": "f", "input": {}} - ], + "content": [{"type": "tool_use", "id": "toolu_a", "name": "f", "input": {}}], }, { "role": "user", @@ -274,9 +272,7 @@ def test_tool_result_list_content_shape_preserved(): }, { "role": "assistant", - "content": [ - {"type": "tool_use", "id": "toolu_b", "name": "f", "input": {}} - ], + "content": [{"type": "tool_use", "id": "toolu_b", "name": "f", "input": {}}], }, { "role": "user", diff --git a/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py b/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py index bfba50fb368..a94c47d73f1 100644 --- a/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py +++ b/tests/unit/llms/anthropic/pass_through/context_management/test_compact.py @@ -196,10 +196,7 @@ def test_select_last_user_question_strips_tool_result_from_mixed_turn(): content = selected[0]["content"] assert isinstance(content, list) assert all(b.get("type") != "tool_result" for b in content) - assert any( - b.get("type") == "text" and b.get("text") == "follow-up question" - for b in content - ) + assert any(b.get("type") == "text" and b.get("text") == "follow-up question" for b in content) def test_select_last_user_question_skips_pure_tool_result_turn(): @@ -428,9 +425,7 @@ def test_client_compaction_block_history_without_context_management(): def test_client_compaction_block_history_no_compaction_returns_none(): - result = apply_client_compaction_block_history( - messages=_simple_messages(), system="base" - ) + result = apply_client_compaction_block_history(messages=_simple_messages(), system="base") assert result is None @@ -515,9 +510,7 @@ async def test_slice_only_no_compaction_block_under_threshold(): async def test_full_summary_path(): """Over threshold: summary call fires, compaction_block and iterations_usage returned.""" messages = _simple_messages() - mock_response = _make_mock_response( - "Condensed history", prompt_tokens=200, completion_tokens=50 - ) + mock_response = _make_mock_response("Condensed history", prompt_tokens=200, completion_tokens=50) with ( patch( @@ -1071,13 +1064,11 @@ async def test_summary_call_does_not_emit_consecutive_user_turns(): ) summary_messages = captured_calls[0]["summary_messages"] - user_indices = [ - idx for idx, msg in enumerate(summary_messages) if msg.get("role") == "user" - ] + user_indices = [idx for idx, msg in enumerate(summary_messages) if msg.get("role") == "user"] # No two adjacent indices. - assert all( - b - a > 1 for a, b in zip(user_indices, user_indices[1:]) - ), f"two consecutive user turns produced: {summary_messages}" + assert all(b - a > 1 for a, b in zip(user_indices, user_indices[1:])), ( + f"two consecutive user turns produced: {summary_messages}" + ) async def test_summary_call_sends_default_max_tokens(): @@ -1160,9 +1151,9 @@ def test_summary_max_tokens_setting_falls_back_for_invalid_values(): "litellm.proxy.proxy_server.general_settings", {"context_management_summary_max_tokens": bad}, ): - assert ( - _read_summary_max_tokens_setting() == COMPACT_SUMMARY_MAX_TOKENS - ), f"expected default for invalid override {bad!r}" + assert _read_summary_max_tokens_setting() == COMPACT_SUMMARY_MAX_TOKENS, ( + f"expected default for invalid override {bad!r}" + ) async def test_summary_call_sends_default_timeout(): @@ -1287,9 +1278,7 @@ async def test_summary_model_denied_when_team_not_in_allowlist(): tools=None, system=None, edit_spec=_EDIT_SPEC_DEFAULT, - user_api_key_auth=_fake_user_api_key_auth( - key_models=["all-proxy-models"], team_models=["gpt-4o"] - ), + user_api_key_auth=_fake_user_api_key_auth(key_models=["all-proxy-models"], team_models=["gpt-4o"]), ) mock_call.assert_not_awaited() @@ -1318,9 +1307,7 @@ async def test_summary_model_allowed_when_in_key_allowlist(): tools=None, system=None, edit_spec=_EDIT_SPEC_DEFAULT, - user_api_key_auth=_fake_user_api_key_auth( - key_models=["claude-haiku-4-5", "gpt-4o"] - ), + user_api_key_auth=_fake_user_api_key_auth(key_models=["claude-haiku-4-5", "gpt-4o"]), ) mock_call.assert_awaited_once() @@ -1575,9 +1562,7 @@ async def test_summary_model_denied_when_key_over_model_budget(): limiter = MagicMock() limiter.is_key_within_model_budget = AsyncMock( - side_effect=litellm.BudgetExceededError( - message="over budget", current_cost=10, max_budget=5 - ) + side_effect=litellm.BudgetExceededError(message="over budget", current_cost=10, max_budget=5) ) with ( @@ -1628,9 +1613,7 @@ async def test_summary_model_denied_when_user_over_model_budget(): limiter = MagicMock() limiter.is_user_within_model_budget = AsyncMock( - side_effect=litellm.BudgetExceededError( - message="over budget", current_cost=10, max_budget=5 - ) + side_effect=litellm.BudgetExceededError(message="over budget", current_cost=10, max_budget=5) ) with ( @@ -1671,9 +1654,7 @@ async def test_summary_model_denied_when_user_over_model_budget(): _PROXY_VirtualKeyModelMaxBudgetLimiter, ) - real_params = inspect.signature( - _PROXY_VirtualKeyModelMaxBudgetLimiter.is_user_within_model_budget - ).parameters + real_params = inspect.signature(_PROXY_VirtualKeyModelMaxBudgetLimiter.is_user_within_model_budget).parameters for kwarg in ("user_id", "user_model_max_budget", "model"): assert kwarg in real_params, f"compact.py passes {kwarg}=, which the limiter no longer accepts" @@ -1695,9 +1676,7 @@ async def test_summary_model_denied_when_end_user_over_model_budget(): limiter = MagicMock() limiter.is_key_within_model_budget = AsyncMock(return_value=True) limiter.is_end_user_within_model_budget = AsyncMock( - side_effect=litellm.BudgetExceededError( - message="over budget", current_cost=10, max_budget=5 - ) + side_effect=litellm.BudgetExceededError(message="over budget", current_cost=10, max_budget=5) ) with ( @@ -2188,9 +2167,7 @@ async def test_model_budget_metadata_propagated_to_summary_call(): parent_litellm_metadata = { "user_api_key": "sk-test", "user_api_key_model_max_budget": {"claude-haiku-4-5": {"budget_limit": 5}}, - "user_api_key_end_user_model_max_budget": { - "claude-haiku-4-5": {"budget_limit": 2} - }, + "user_api_key_end_user_model_max_budget": {"claude-haiku-4-5": {"budget_limit": 2}}, } with ( @@ -2215,12 +2192,8 @@ async def test_model_budget_metadata_propagated_to_summary_call(): ) propagated = mock_call.call_args.kwargs["metadata"] - assert propagated["user_api_key_model_max_budget"] == { - "claude-haiku-4-5": {"budget_limit": 5} - } - assert propagated["user_api_key_end_user_model_max_budget"] == { - "claude-haiku-4-5": {"budget_limit": 2} - } + assert propagated["user_api_key_model_max_budget"] == {"claude-haiku-4-5": {"budget_limit": 5}} + assert propagated["user_api_key_end_user_model_max_budget"] == {"claude-haiku-4-5": {"budget_limit": 2}} async def test_summary_call_propagates_allowed_model_region(): @@ -2692,9 +2665,7 @@ def test_endpoint_returns_anthropic_400_on_context_management_error(): mock_proxy_server.version = "test" with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}): - with patch( - "litellm.proxy.anthropic_endpoints.endpoints.ProxyBaseLLMRequestProcessing" - ) as mock_cls: + with patch("litellm.proxy.anthropic_endpoints.endpoints.ProxyBaseLLMRequestProcessing") as mock_cls: mock_instance = MagicMock() mock_instance.base_process_llm_request = AsyncMock( side_effect=AnthropicContextManagementError( @@ -2753,9 +2724,7 @@ def test_endpoint_runs_failure_hook_on_500_context_management_error(): mock_proxy_server.version = "test" with patch.dict(sys.modules, {"litellm.proxy.proxy_server": mock_proxy_server}): - with patch( - "litellm.proxy.anthropic_endpoints.endpoints.ProxyBaseLLMRequestProcessing" - ) as mock_cls: + with patch("litellm.proxy.anthropic_endpoints.endpoints.ProxyBaseLLMRequestProcessing") as mock_cls: mock_instance = MagicMock() mock_instance.base_process_llm_request = AsyncMock( side_effect=AnthropicContextManagementError( diff --git a/tests/unit/llms/anthropic/pass_through/context_management/test_dispatcher.py b/tests/unit/llms/anthropic/pass_through/context_management/test_dispatcher.py index 5943661683a..e28d2ac5e65 100644 --- a/tests/unit/llms/anthropic/pass_through/context_management/test_dispatcher.py +++ b/tests/unit/llms/anthropic/pass_through/context_management/test_dispatcher.py @@ -50,9 +50,7 @@ async def test_unknown_edit_type_is_noop(): messages=messages, tools=None, system=None, - context_management_spec={ - "edits": [{"type": "totally_not_a_real_edit_20999999"}] - }, + context_management_spec={"edits": [{"type": "totally_not_a_real_edit_20999999"}]}, ) assert result.applied_edits == [] assert result.messages == messages diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_advisor_integration.py b/tests/unit/llms/anthropic/pass_through/messages/test_advisor_integration.py index 57d45854130..f5f30fb86c1 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_advisor_integration.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_advisor_integration.py @@ -39,9 +39,7 @@ def _text_resp(text: str, model: str = "gpt-4o-mini") -> Dict: } -def _advisor_call_resp( - question: str = "How do I approach this?", tool_id: str = "tid_01" -) -> Dict: +def _advisor_call_resp(question: str = "How do I approach this?", tool_id: str = "tid_01") -> Dict: return { "id": "msg_int_test", "type": "message", @@ -106,14 +104,10 @@ async def test_full_dispatch_interceptor_fires_and_loop_completes(): assert isinstance(result, dict) content = result.get("content", []) text_blocks = [b for b in content if b.get("type") == "text"] - advisor_uses = [ - b for b in content if b.get("type") == "tool_use" and b.get("name") == "advisor" - ] + advisor_uses = [b for b in content if b.get("type") == "tool_use" and b.get("name") == "advisor"] assert len(text_blocks) >= 1, "Final response must have text" - assert ( - len(advisor_uses) == 0 - ), "No advisor tool_use blocks must appear in final output" + assert len(advisor_uses) == 0, "No advisor tool_use blocks must appear in final output" # --------------------------------------------------------------------------- @@ -221,9 +215,7 @@ async def test_named_params_forwarded_into_advisor_executor_subcall(): captured_executor_kwargs: Dict = {} - async def mock_handler( - model, messages, tools, stream, max_tokens, custom_llm_provider, **kwargs - ): + async def mock_handler(model, messages, tools, stream, max_tokens, custom_llm_provider, **kwargs): # First call is the executor sub-call (returns advisor tool_use). # Capture its kwargs so we can assert the forwarded params. if not captured_executor_kwargs: @@ -267,8 +259,7 @@ async def test_named_params_forwarded_into_advisor_executor_subcall(): ) assert captured_executor_kwargs["thinking"] == {"type": "adaptive"}, ( - "thinking must be forwarded into executor sub-call — see " - "anthropic_messages.handler interceptor invocation." + "thinking must be forwarded into executor sub-call — see anthropic_messages.handler interceptor invocation." ) # The advisor enriches metadata with `advisor_sub_call` / `parent_request_id`, # but the original caller fields must survive into the executor sub-call. @@ -304,9 +295,7 @@ async def test_pre_request_hook_override_does_not_collide_with_explicit_kwargs() captured: Dict = {} - async def mock_handler( - model, messages, tools, stream, max_tokens, custom_llm_provider, **kwargs - ): + async def mock_handler(model, messages, tools, stream, max_tokens, custom_llm_provider, **kwargs): if not captured: captured.update( { @@ -320,9 +309,7 @@ async def test_pre_request_hook_override_does_not_collide_with_explicit_kwargs() return _text_resp("Some advice.", model="claude-opus-4-6") return _text_resp("Final answer.") - async def fake_pre_request_hooks( - model, messages, tools, stream, custom_llm_provider, **hook_kwargs - ): + async def fake_pre_request_hooks(model, messages, tools, stream, custom_llm_provider, **hook_kwargs): # Simulate a CustomLogger.async_pre_request_hook that overrides several # named params on its way through. Without the request_kwargs.pop() # extraction in handler.py, these would collide with the explicit diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_agentic_streaming_iterator.py b/tests/unit/llms/anthropic/pass_through/messages/test_agentic_streaming_iterator.py index 16244db04a3..56fa3015d44 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_agentic_streaming_iterator.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_agentic_streaming_iterator.py @@ -73,9 +73,7 @@ def _build_simple_text_stream() -> List[bytes]: }, ) ) - chunks.append( - _sse_event("content_block_stop", {"type": "content_block_stop", "index": 0}) - ) + chunks.append(_sse_event("content_block_stop", {"type": "content_block_stop", "index": 0})) chunks.append( _sse_event( "message_delta", @@ -148,9 +146,7 @@ def _build_tool_use_stream() -> List[bytes]: }, ) ) - chunks.append( - _sse_event("content_block_stop", {"type": "content_block_stop", "index": 0}) - ) + chunks.append(_sse_event("content_block_stop", {"type": "content_block_stop", "index": 0})) # tool_use block chunks.append( _sse_event( @@ -190,9 +186,7 @@ def _build_tool_use_stream() -> List[bytes]: }, ) ) - chunks.append( - _sse_event("content_block_stop", {"type": "content_block_stop", "index": 1}) - ) + chunks.append(_sse_event("content_block_stop", {"type": "content_block_stop", "index": 1})) chunks.append( _sse_event( "message_delta", @@ -284,9 +278,7 @@ def _build_hold_back_iterator( class TestParseSSEEvents: def test_should_parse_single_event(self): - raw = _sse_event( - "message_start", {"type": "message_start", "message": {"id": "1"}} - ) + raw = _sse_event("message_start", {"type": "message_start", "message": {"id": "1"}}) events = _parse_sse_events(raw) assert len(events) == 1 assert events[0][0] == "message_start" @@ -457,9 +449,7 @@ class TestHandleMessageDelta: class TestRebuildAnthropicResponse: def test_should_rebuild_simple_text_response(self): raw_bytes = _build_simple_text_stream() - result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse( - raw_bytes - ) + result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(raw_bytes) assert result is not None assert result["id"] == "msg_123" assert result["model"] == "claude-sonnet-4-20250514" @@ -472,9 +462,7 @@ class TestRebuildAnthropicResponse: def test_should_rebuild_tool_use_response(self): raw_bytes = _build_tool_use_stream() - result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse( - raw_bytes - ) + result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(raw_bytes) assert result is not None assert result["id"] == "msg_tool_456" assert result["stop_reason"] == "tool_use" @@ -502,23 +490,17 @@ class TestRebuildAnthropicResponse: }, ) ] - result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse( - raw_bytes - ) + result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(raw_bytes) assert result is None def test_should_handle_empty_bytes(self): - result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse( - [] - ) + result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse([]) assert result is None def test_should_handle_multi_event_chunks(self): """When multiple SSE events arrive in a single bytes chunk.""" combined = b"".join(_build_simple_text_stream()) - result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse( - [combined] - ) + result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse([combined]) assert result is not None assert result["content"][0]["text"] == "Hello, world!" @@ -550,9 +532,7 @@ class TestRebuildAnthropicResponse: ), _sse_event("message_stop", {"type": "message_stop"}), ] - result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse( - raw_bytes - ) + result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(raw_bytes) assert result is not None assert result["usage"]["cache_creation_input_tokens"] == 50 assert result["usage"]["cache_read_input_tokens"] == 30 @@ -593,9 +573,7 @@ class TestRebuildAnthropicResponse: ), _sse_event("message_stop", {"type": "message_stop"}), ] - result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse( - raw_bytes - ) + result = AgenticAnthropicStreamingIterator._rebuild_anthropic_response_from_sse(raw_bytes) assert result is not None assert result["content"][0]["type"] == "redacted_thinking" @@ -722,9 +700,7 @@ class TestAgenticStreamingIteratorPhase2: } mock_handler = MagicMock() - mock_handler._call_agentic_completion_hooks = AsyncMock( - return_value=fake_response - ) + mock_handler._call_agentic_completion_hooks = AsyncMock(return_value=fake_response) iterator = AgenticAnthropicStreamingIterator( completion_stream=mock_stream, @@ -757,9 +733,7 @@ class TestAgenticStreamingIteratorErrorHandling: mock_stream = MockAsyncStream(chunks) mock_handler = MagicMock() - mock_handler._call_agentic_completion_hooks = AsyncMock( - side_effect=RuntimeError("hook exploded") - ) + mock_handler._call_agentic_completion_hooks = AsyncMock(side_effect=RuntimeError("hook exploded")) mock_logging = MagicMock() mock_logging.litellm_call_id = "test_call_123" diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_transformation.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_transformation.py new file mode 100644 index 00000000000..0481f7785cd --- /dev/null +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_experimental_pass_through_messages_transformation.py @@ -0,0 +1,118 @@ +from pathlib import Path +from typing import Final + +import pytest + +import litellm +from litellm.llms.anthropic.pass_through.messages.transformation import ( + AnthropicMessagesConfig, +) +from litellm.llms.anthropic.wif import get_anthropic_wif_token +from litellm.llms.minimax.messages.transformation import MinimaxMessagesConfig +from litellm.llms.tencent.messages.transformation import TencentAnthropicMessagesConfig +from tests.unit.llms.anthropic.test_anthropic_wif import ( + ScriptedPoster, + make_engine, + token_response, + write_token_file, +) + +_WIF_PARAMS: Final[dict] = { + "anthropic_federation_rule_id": "fdrl_abc123", + "anthropic_organization_id": "org-uuid-1", + "anthropic_identity_token_file": "/var/run/secrets/identity-token", +} + + +def test_workload_identity_allowed_for_anthropic() -> None: + assert AnthropicMessagesConfig()._allows_workload_identity is True + + +def test_workload_identity_blocked_for_minimax() -> None: + assert MinimaxMessagesConfig()._allows_workload_identity is False + + +def test_workload_identity_blocked_for_tencent() -> None: + assert TencentAnthropicMessagesConfig()._allows_workload_identity is False + + +def test_minimax_validate_environment_never_attaches_anthropic_wif_credential( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """Regression test: before the fix, an Anthropic-WIF-configured proxy would mint a real + Anthropic federation token inside MiniMax's inherited validate_anthropic_messages_environment + and send it as the Authorization header on the MiniMax-routed request. With no MiniMax + credential of its own the deployment must fail closed on the missing key instead.""" + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_prod") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-prod-uuid") + monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) + monkeypatch.delenv("MINIMAX_API_KEY", raising=False) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + litellm_params = {"anthropic_identity_token_file": str(token_file)} + + with pytest.raises(litellm.AuthenticationError, match="Missing Anthropic API Key"): + MinimaxMessagesConfig().validate_anthropic_messages_environment( + headers={}, + model="MiniMax-M2.1", + messages=[], + optional_params={}, + litellm_params=litellm_params, + api_key=None, + api_base="https://api.minimax.io/anthropic", + ) + + +def test_tencent_validate_environment_never_attaches_anthropic_wif_credential( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_prod") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-prod-uuid") + monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) + monkeypatch.delenv("TENCENT_API_KEY", raising=False) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + litellm_params = {"anthropic_identity_token_file": str(token_file)} + + monkeypatch.setattr(litellm, "api_key", None) + with pytest.raises(litellm.AuthenticationError, match="Missing Anthropic API Key"): + TencentAnthropicMessagesConfig().validate_anthropic_messages_environment( + headers={}, + model="deepseek-v4-pro", + messages=[], + optional_params={}, + litellm_params=litellm_params, + api_key=None, + api_base="https://tokenhub-intl.tencentcloudmaas.com", + ) + + +def test_wif_token_exchange_reaches_only_anthropic_not_minimax_or_tencent( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch +) -> None: + """get_anthropic_wif_token's engine parameter is the only DI seam in the WIF minting chain; + validate_anthropic_messages_environment always uses the module's default engine, so this + drives that seam directly with the exact litellm_params AnthropicModelInfo.get_auth_header + would receive from each config, proving MiniMax/Tencent never reach the token endpoint even + when a mint would otherwise succeed.""" + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_prod") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-prod-uuid") + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + litellm_params = {"anthropic_identity_token_file": str(token_file)} + poster = ScriptedPoster([token_response("sk-ant-oat01-canary")]) + engine = make_engine(poster) + + minted: Final = get_anthropic_wif_token( + litellm_params, + "https://api.anthropic.com", + "claude-sonnet-4-5", + engine, + ) + assert minted == "sk-ant-oat01-canary" + assert len(poster.requests) == 1 + + for config in (MinimaxMessagesConfig(), TencentAnthropicMessagesConfig()): + assert config._allows_workload_identity is False + + assert len(poster.requests) == 1 diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_speed.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_speed.py index 609a9fd73a5..45c6ebaa39b 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_speed.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_speed.py @@ -93,13 +93,11 @@ def test_messages_drops_speed_for_vertex_opus_with_drop_params(monkeypatch): """Regression: a vertex_ai Opus passthrough must drop ``speed`` even though the prefix-stripped model id maps to a fast-mode-capable direct-Anthropic entry.""" monkeypatch.setattr(litellm, "drop_params", True) - optional_params = ( - AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( - params={"max_tokens": 1024, "speed": "fast"}, - model="claude-opus-4-8", - drop_params=False, - custom_llm_provider="vertex_ai", - ) + optional_params = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( + params={"max_tokens": 1024, "speed": "fast"}, + model="claude-opus-4-8", + drop_params=False, + custom_llm_provider="vertex_ai", ) assert "speed" not in optional_params diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_structured_outputs.py b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_structured_outputs.py index d1e17590224..670e3de5110 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_structured_outputs.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_anthropic_messages_structured_outputs.py @@ -13,9 +13,7 @@ def test_output_format_supported_and_transforms_correctly(): config = AnthropicMessagesConfig() # 1. Verify it's in supported parameters - supported_params = config.get_supported_anthropic_messages_params( - "claude-sonnet-4-5" - ) + supported_params = config.get_supported_anthropic_messages_params("claude-sonnet-4-5") assert "output_format" in supported_params # 2. Verify transformation preserves output_format and adds beta header diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_content_after_stop_reason.py b/tests/unit/llms/anthropic/pass_through/messages/test_content_after_stop_reason.py index 7154b10aaca..f57a872c82a 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_content_after_stop_reason.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_content_after_stop_reason.py @@ -30,26 +30,14 @@ class MockCompletionStreamWithContentAfterStopReason: self.responses = [ # Initial text content ModelResponseStream( - choices=[ - StreamingChoices( - delta=Delta(content="Hello"), index=0, finish_reason=None - ) - ], + choices=[StreamingChoices(delta=Delta(content="Hello"), index=0, finish_reason=None)], ), ModelResponseStream( - choices=[ - StreamingChoices( - delta=Delta(content=" world"), index=0, finish_reason=None - ) - ], + choices=[StreamingChoices(delta=Delta(content=" world"), index=0, finish_reason=None)], ), # Message delta with stop_reason AND usage (this is how it actually comes from the API) ModelResponseStream( - choices=[ - StreamingChoices( - delta=Delta(content=""), index=0, finish_reason="stop" - ) - ], + choices=[StreamingChoices(delta=Delta(content=""), index=0, finish_reason="stop")], usage=Usage(prompt_tokens=230, completion_tokens=65, total_tokens=295), ), # Additional content after the stop_reason - this simulates the scenario @@ -118,9 +106,9 @@ def test_anthropic_stream_wrapper_content_after_stop_reason(): print(f"Expected chunk types: {expected_types}") # Verify we have the expected number of chunks - assert len(chunk_types) >= len( - expected_types - ), f"Expected at least {len(expected_types)} chunks, got {len(chunk_types)}" + assert len(chunk_types) >= len(expected_types), ( + f"Expected at least {len(expected_types)} chunks, got {len(chunk_types)}" + ) # Verify key chunk types are present assert "message_start" in chunk_types @@ -143,15 +131,9 @@ def test_anthropic_stream_wrapper_content_after_stop_reason(): delta = message_delta_chunk.get("delta", {}) usage = message_delta_chunk.get("usage", {}) - assert ( - delta.get("stop_reason") == "end_turn" - ), f"Expected stop_reason 'end_turn', got {delta.get('stop_reason')}" - assert ( - usage.get("input_tokens") == 230 - ), f"Expected input_tokens 230, got {usage.get('input_tokens')}" - assert ( - usage.get("output_tokens") == 65 - ), f"Expected output_tokens 65, got {usage.get('output_tokens')}" + assert delta.get("stop_reason") == "end_turn", f"Expected stop_reason 'end_turn', got {delta.get('stop_reason')}" + assert usage.get("input_tokens") == 230, f"Expected input_tokens 230, got {usage.get('input_tokens')}" + assert usage.get("output_tokens") == 65, f"Expected output_tokens 65, got {usage.get('output_tokens')}" # Verify content_block_stop comes before message_delta content_block_stop_index = None @@ -165,9 +147,7 @@ def test_anthropic_stream_wrapper_content_after_stop_reason(): assert content_block_stop_index is not None, "content_block_stop not found" assert message_delta_index is not None, "message_delta not found" - assert ( - content_block_stop_index < message_delta_index - ), "content_block_stop should come before message_delta" + assert content_block_stop_index < message_delta_index, "content_block_stop should come before message_delta" @pytest.mark.asyncio @@ -210,15 +190,9 @@ async def test_async_anthropic_stream_wrapper_content_after_stop_reason(): delta = message_delta_chunk.get("delta", {}) usage = message_delta_chunk.get("usage", {}) - assert ( - delta.get("stop_reason") == "end_turn" - ), f"Expected stop_reason 'end_turn', got {delta.get('stop_reason')}" - assert ( - usage.get("input_tokens") == 230 - ), f"Expected input_tokens 230, got {usage.get('input_tokens')}" - assert ( - usage.get("output_tokens") == 65 - ), f"Expected output_tokens 65, got {usage.get('output_tokens')}" + assert delta.get("stop_reason") == "end_turn", f"Expected stop_reason 'end_turn', got {delta.get('stop_reason')}" + assert usage.get("input_tokens") == 230, f"Expected input_tokens 230, got {usage.get('input_tokens')}" + assert usage.get("output_tokens") == 65, f"Expected output_tokens 65, got {usage.get('output_tokens')}" def test_usage_merging_behavior(): @@ -234,18 +208,10 @@ def test_usage_merging_behavior(): for chunk in wrapper: chunks.append(chunk) # If this is a message_delta with stop_reason, verify it has usage - if ( - chunk.get("type") == "message_delta" - and chunk.get("delta", {}).get("stop_reason") is not None - ): - + if chunk.get("type") == "message_delta" and chunk.get("delta", {}).get("stop_reason") is not None: usage = chunk.get("usage", {}) - assert ( - usage.get("input_tokens") is not None - ), "Usage should be merged with stop_reason chunk" - assert ( - usage.get("output_tokens") is not None - ), "Usage should be merged with stop_reason chunk" + assert usage.get("input_tokens") is not None, "Usage should be merged with stop_reason chunk" + assert usage.get("output_tokens") is not None, "Usage should be merged with stop_reason chunk" break @@ -273,12 +239,8 @@ def test_sse_wrapper_with_content_after_stop_reason(): lines = chunk_str.split("\n") # Should have event and data lines - assert any( - line.startswith("event: ") for line in lines - ), f"Missing event line in: {chunk_str}" - assert any( - line.startswith("data: ") for line in lines - ), f"Missing data line in: {chunk_str}" + assert any(line.startswith("event: ") for line in lines), f"Missing event line in: {chunk_str}" + assert any(line.startswith("data: ") for line in lines), f"Missing data line in: {chunk_str}" @pytest.mark.asyncio @@ -306,12 +268,8 @@ async def test_async_sse_wrapper_with_content_after_stop_reason(): lines = chunk_str.split("\n") # Should have event and data lines - assert any( - line.startswith("event: ") for line in lines - ), f"Missing event line in: {chunk_str}" - assert any( - line.startswith("data: ") for line in lines - ), f"Missing data line in: {chunk_str}" + assert any(line.startswith("event: ") for line in lines), f"Missing event line in: {chunk_str}" + assert any(line.startswith("data: ") for line in lines), f"Missing data line in: {chunk_str}" if __name__ == "__main__": diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_parallel_tool_calls.py b/tests/unit/llms/anthropic/pass_through/messages/test_parallel_tool_calls.py index 45e39a572c5..229eff4275a 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_parallel_tool_calls.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_parallel_tool_calls.py @@ -53,9 +53,7 @@ def construct_text_chunk(text: str) -> ModelResponseStream: ) -def construct_split_tool_call( - id: str, function_name: str, function_arg_parts: List[str] -) -> List[ModelResponseStream]: +def construct_split_tool_call(id: str, function_name: str, function_arg_parts: List[str]) -> List[ModelResponseStream]: return [ # https://platform.openai.com/docs/guides/function-calling#streaming ModelResponseStream( @@ -144,10 +142,7 @@ def test_anthropic_stream_wrapper_single_tool_call(): get_weather_calls = 0 for chunk in chunks: - if ( - chunk.get("type") == "content_block_start" - and chunk["content_block"]["type"] == "tool_use" - ): + if chunk.get("type") == "content_block_start" and chunk["content_block"]["type"] == "tool_use": if chunk["content_block"]["name"] == "get_weather": get_weather_calls += 1 @@ -203,10 +198,7 @@ def test_anthropic_stream_wrapper_back_to_back_tool_calls(): get_weather_calls = 0 for chunk in chunks: - if ( - chunk.get("type") == "content_block_start" - and chunk["content_block"]["type"] == "tool_use" - ): + if chunk.get("type") == "content_block_start" and chunk["content_block"]["type"] == "tool_use": if chunk["content_block"]["name"] == "get_weather": get_weather_calls += 1 @@ -218,9 +210,7 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text(): *construct_split_tool_call("tooluse_foo", "get_weather", ['{"city":', '"NY"}']), construct_text_chunk("The weather is nice today."), *construct_split_tool_call("tooluse_bar", "get_weather", ['{"city":', '"SF"}']), - *construct_split_tool_call( - "tooluse_bar", "get_weather", ['{"city":', '"CHI"}'] - ), + *construct_split_tool_call("tooluse_bar", "get_weather", ['{"city":', '"CHI"}']), construct_text_chunk("The weather is not so nice today."), ModelResponseStream( choices=[ @@ -280,8 +270,7 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text(): text_deltas = [ chunk["delta"]["text"] for chunk in chunks - if chunk.get("type") == "content_block_delta" - and chunk["delta"].get("type") == "text_delta" + if chunk.get("type") == "content_block_delta" and chunk["delta"].get("type") == "text_delta" ] assert text_deltas == [ "The weather is nice today.", @@ -291,10 +280,7 @@ def test_anthropic_stream_wrapper_interleaved_tool_calls_and_text(): get_weather_calls = 0 for chunk in chunks: - if ( - chunk.get("type") == "content_block_start" - and chunk["content_block"]["type"] == "tool_use" - ): + if chunk.get("type") == "content_block_start" and chunk["content_block"]["type"] == "tool_use": if chunk["content_block"]["name"] == "get_weather": get_weather_calls += 1 diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_reasoning_auto_summary_messages.py b/tests/unit/llms/anthropic/pass_through/messages/test_reasoning_auto_summary_messages.py index 42c7814e42e..cc951421b10 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_reasoning_auto_summary_messages.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_reasoning_auto_summary_messages.py @@ -71,9 +71,7 @@ class TestReasoningAutoSummaryMessages: def test_adaptive_thinking_gets_display_summarized(self): """reasoning_auto_summary=True + thinking.type='adaptive' -> display='summarized'.""" with patch.object(litellm, "reasoning_auto_summary", True): - params = _call_handler_and_capture_optional_params( - thinking={"type": "adaptive", "budget_tokens": 5000} - ) + params = _call_handler_and_capture_optional_params(thinking={"type": "adaptive", "budget_tokens": 5000}) thinking = params.get("thinking", {}) assert thinking.get("display") == "summarized" assert thinking.get("type") == "adaptive" @@ -82,9 +80,7 @@ class TestReasoningAutoSummaryMessages: def test_enabled_thinking_gets_display_summarized(self): """reasoning_auto_summary=True + thinking.type='enabled' -> display='summarized'.""" with patch.object(litellm, "reasoning_auto_summary", True): - params = _call_handler_and_capture_optional_params( - thinking={"type": "enabled", "budget_tokens": 10000} - ) + params = _call_handler_and_capture_optional_params(thinking={"type": "enabled", "budget_tokens": 10000}) thinking = params.get("thinking", {}) assert thinking.get("display") == "summarized" assert thinking.get("type") == "enabled" @@ -92,18 +88,14 @@ class TestReasoningAutoSummaryMessages: def test_disabled_thinking_no_display(self): """reasoning_auto_summary=True + thinking.type='disabled' -> display NOT set.""" with patch.object(litellm, "reasoning_auto_summary", True): - params = _call_handler_and_capture_optional_params( - thinking={"type": "disabled"} - ) + params = _call_handler_and_capture_optional_params(thinking={"type": "disabled"}) thinking = params.get("thinking", {}) assert "display" not in thinking def test_no_injection_when_flag_false(self): """reasoning_auto_summary=False + active thinking -> display NOT set.""" with patch.object(litellm, "reasoning_auto_summary", False): - params = _call_handler_and_capture_optional_params( - thinking={"type": "enabled", "budget_tokens": 10000} - ) + params = _call_handler_and_capture_optional_params(thinking={"type": "enabled", "budget_tokens": 10000}) thinking = params.get("thinking", {}) assert "display" not in thinking @@ -117,12 +109,11 @@ class TestReasoningAutoSummaryMessages: def test_env_var_enables_auto_summary(self): """LITELLM_REASONING_AUTO_SUMMARY=true env var enables the feature.""" - with patch.object(litellm, "reasoning_auto_summary", False), patch.dict( - os.environ, {"LITELLM_REASONING_AUTO_SUMMARY": "true"} + with ( + patch.object(litellm, "reasoning_auto_summary", False), + patch.dict(os.environ, {"LITELLM_REASONING_AUTO_SUMMARY": "true"}), ): - params = _call_handler_and_capture_optional_params( - thinking={"type": "adaptive", "budget_tokens": 5000} - ) + params = _call_handler_and_capture_optional_params(thinking={"type": "adaptive", "budget_tokens": 5000}) thinking = params.get("thinking", {}) assert thinking.get("display") == "summarized" diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py b/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py index dc4da8198cd..dd744ca66a1 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_request_optional_param_utils.py @@ -23,11 +23,7 @@ def test_optional_param_filtering_unchanged(): "not_a_real_param": "drop me", # invalid key dropped "stream": True, } - result = ( - AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( - params - ) - ) + result = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param(params) assert result == {"temperature": 0.5, "tools": [{"name": "x"}], "stream": True} assert "top_p" not in result assert "not_a_real_param" not in result @@ -37,9 +33,7 @@ def test_valid_keys_are_memoized(): _anthropic_messages_optional_param_keys.cache_clear() first = _anthropic_messages_optional_param_keys() for _ in range(50): - AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( - {"temperature": 0.1} - ) + AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param({"temperature": 0.1}) info = _anthropic_messages_optional_param_keys.cache_info() # Resolved exactly once despite many calls. assert info.misses == 1 @@ -51,23 +45,16 @@ def test_valid_keys_are_memoized(): def test_empty_params(): - assert ( - AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( - {} - ) - == {} - ) + assert AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param({}) == {} def test_drop_params_strips_speed_for_unsupported_model(): original = litellm.drop_params litellm.drop_params = True try: - result = ( - AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( - params={"speed": "fast", "temperature": 0.5}, - model="claude-sonnet-4-6", - ) + result = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( + params={"speed": "fast", "temperature": 0.5}, + model="claude-sonnet-4-6", ) finally: litellm.drop_params = original @@ -80,11 +67,9 @@ def test_drop_params_keeps_speed_for_supporting_model(): original = litellm.drop_params litellm.drop_params = True try: - result = ( - AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( - params={"speed": "fast"}, - model="claude-opus-4-6", - ) + result = AnthropicMessagesRequestUtils.get_requested_anthropic_messages_optional_param( + params={"speed": "fast"}, + model="claude-opus-4-6", ) finally: litellm.drop_params = original diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_sse_wrapper.py b/tests/unit/llms/anthropic/pass_through/messages/test_sse_wrapper.py index 92f2dce7331..d22478e16ef 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_sse_wrapper.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_sse_wrapper.py @@ -1,4 +1,3 @@ - import pytest from fastapi.testclient import TestClient @@ -14,25 +13,13 @@ class MockCompletionStream: def __init__(self): self.responses = [ ModelResponseStream( - choices=[ - StreamingChoices( - delta=Delta(content="Hello"), index=0, finish_reason=None - ) - ], + choices=[StreamingChoices(delta=Delta(content="Hello"), index=0, finish_reason=None)], ), ModelResponseStream( - choices=[ - StreamingChoices( - delta=Delta(content=" World"), index=0, finish_reason=None - ) - ], + choices=[StreamingChoices(delta=Delta(content=" World"), index=0, finish_reason=None)], ), ModelResponseStream( - choices=[ - StreamingChoices( - delta=Delta(content=""), index=0, finish_reason="stop" - ) - ], + choices=[StreamingChoices(delta=Delta(content=""), index=0, finish_reason="stop")], ), ] self.index = 0 @@ -50,9 +37,7 @@ class MockCompletionStream: def test_anthropic_sse_wrapper_format(): """Test that the SSE wrapper produces proper event and data formatting""" - wrapper = AnthropicStreamWrapper( - completion_stream=MockCompletionStream(), model="claude-3" - ) + wrapper = AnthropicStreamWrapper(completion_stream=MockCompletionStream(), model="claude-3") # Get the first chunk from the SSE wrapper first_chunk = next(wrapper.anthropic_sse_wrapper()) @@ -73,9 +58,7 @@ def test_anthropic_sse_wrapper_format(): def test_anthropic_sse_wrapper_event_types(): """Test that different chunk types produce correct event types""" - wrapper = AnthropicStreamWrapper( - completion_stream=MockCompletionStream(), model="claude-3" - ) + wrapper = AnthropicStreamWrapper(completion_stream=MockCompletionStream(), model="claude-3") chunks = [] for chunk in wrapper.anthropic_sse_wrapper(): @@ -104,18 +87,10 @@ async def test_async_anthropic_sse_wrapper(): def __init__(self): self.responses = [ ModelResponseStream( - choices=[ - StreamingChoices( - delta=Delta(content="Hello"), index=0, finish_reason=None - ) - ], + choices=[StreamingChoices(delta=Delta(content="Hello"), index=0, finish_reason=None)], ), ModelResponseStream( - choices=[ - StreamingChoices( - delta=Delta(content=" World"), index=0, finish_reason=None - ) - ], + choices=[StreamingChoices(delta=Delta(content=" World"), index=0, finish_reason=None)], ), ] self.index = 0 @@ -130,9 +105,7 @@ async def test_async_anthropic_sse_wrapper(): self.index += 1 return response - wrapper = AnthropicStreamWrapper( - completion_stream=AsyncMockCompletionStream(), model="claude-3" - ) + wrapper = AnthropicStreamWrapper(completion_stream=AsyncMockCompletionStream(), model="claude-3") # Get the first chunk from the async SSE wrapper first_chunk = None diff --git a/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py b/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py index 39c5b8048c8..54e25da6ab9 100644 --- a/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py +++ b/tests/unit/llms/anthropic/pass_through/messages/test_streaming_iterator.py @@ -167,7 +167,7 @@ async def test_async_sse_wrapper_treats_message_stop_bytes_as_complete(): def test_is_message_stop_chunk(): assert _is_message_stop_chunk({"type": "message_stop"}) is True assert _is_message_stop_chunk({"type": "message_delta"}) is False - assert _is_message_stop_chunk(b'event: message_stop\ndata: {}\n\n') is True + assert _is_message_stop_chunk(b"event: message_stop\ndata: {}\n\n") is True assert _is_message_stop_chunk(b"raw-bytes") is False assert _is_message_stop_chunk("message_stop") is False @@ -198,7 +198,7 @@ def test_is_message_stop_chunk_ignores_substring_in_payload(): not be treated as a terminal stop event. """ delta_frame_with_substring = ( - b'event: content_block_delta\n' + b"event: content_block_delta\n" b'data: {"type": "content_block_delta", "delta": ' b'{"type": "input_json_delta", "partial_json": "\\"message_stop\\""}}\n\n' ) @@ -302,10 +302,11 @@ async def test_async_sse_wrapper_emits_error_when_bytes_stream_only_mentions_mes payload text contains `message_stop` (but never emits the actual `event: message_stop` frame) must still be flagged as incomplete. """ + async def _byte_stream(): yield b'event: message_start\ndata: {"type": "message_start"}\n\n' yield ( - b'event: content_block_delta\n' + b"event: content_block_delta\n" b'data: {"type": "content_block_delta", "delta": ' b'{"type": "input_json_delta", "partial_json": "\\"message_stop\\""}}\n\n' ) diff --git a/tests/unit/llms/anthropic/test_anthropic_common_utils.py b/tests/unit/llms/anthropic/test_anthropic_common_utils.py index 68e2e9650d5..068397da323 100644 --- a/tests/unit/llms/anthropic/test_anthropic_common_utils.py +++ b/tests/unit/llms/anthropic/test_anthropic_common_utils.py @@ -13,14 +13,18 @@ Verifies that: import json import os import sys +import threading from types import SimpleNamespace from typing import Final from unittest.mock import patch +import httpx import pytest sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))) +from litellm.proxy._types import SpecialHeaders # noqa: E402 # sys.path must be patched before importing litellm + # Fake tokens for testing (not real secrets) FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef" FAKE_REGULAR_KEY = "sk-ant-api03-regular-key-for-testing-123456789" @@ -1106,14 +1110,19 @@ class TestGetAuthHeader: assert result is None def test_oauth_token_uses_bearer_not_x_api_key(self): - """OAuth token (sk-ant-oat*) should return Authorization: Bearer, not x-api-key.""" + """OAuth token (sk-ant-oat*) should return Authorization: Bearer with the + mandatory oauth beta, not x-api-key.""" from litellm.llms.anthropic.common_utils import AnthropicModelInfo result = AnthropicModelInfo.get_auth_header(api_key=FAKE_OAUTH_TOKEN) - assert result == {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"} + assert result == { + "authorization": f"Bearer {FAKE_OAUTH_TOKEN}", + "anthropic-beta": "oauth-2025-04-20", + } def test_oauth_token_from_env_uses_bearer(self): - """OAuth token in ANTHROPIC_API_KEY env var should return Authorization: Bearer.""" + """OAuth token in ANTHROPIC_API_KEY env var should return Authorization: Bearer + with the mandatory oauth beta.""" from unittest.mock import patch as mock_patch from litellm.llms.anthropic.common_utils import AnthropicModelInfo @@ -1124,7 +1133,10 @@ class TestGetAuthHeader: clear=True, ): result = AnthropicModelInfo.get_auth_header() - assert result == {"authorization": f"Bearer {FAKE_OAUTH_TOKEN}"} + assert result == { + "authorization": f"Bearer {FAKE_OAUTH_TOKEN}", + "anthropic-beta": "oauth-2025-04-20", + } def test_custom_api_base_get_auth_header_uses_bearer(self): """Non-standard API key and custom api_base returns Bearer when use_bearer_for_custom_base=True.""" @@ -2255,6 +2267,1611 @@ def test_create_anthropic_model_list_response_empty(): assert response["last_id"] is None +# --------------------------------------------------------------------------- # +# Workload identity federation wiring (issue #28607) +# --------------------------------------------------------------------------- # + +FAKE_MINTED_TOKEN = "sk-ant-oat01-wif-minted-token-for-testing-abc123" + +ANTHROPIC_ENV_VARS = ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + "ANTHROPIC_FEDERATION_RULE_ID", + "ANTHROPIC_ORGANIZATION_ID", + "ANTHROPIC_SERVICE_ACCOUNT_ID", + "ANTHROPIC_FEDERATION_WORKSPACE_ID", + "ANTHROPIC_IDENTITY_TOKEN_FILE", + "ANTHROPIC_IDENTITY_TOKEN", + "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", +) + +WIF_ENV = { + "ANTHROPIC_FEDERATION_RULE_ID": "fdrl_wire1", + "ANTHROPIC_ORGANIZATION_ID": "org-wire-1", + "ANTHROPIC_IDENTITY_TOKEN": "inline-wire-jwt", +} + +PROXY_CREDENTIAL_HEADER_NAMES = sorted(SpecialHeaders.litellm_credential_header_names()) + + +class RecordingPoster: + def __init__(self, response): + self.requests = [] + self.thread_ids = [] + self._response = response + + def post(self, url, *, content, headers, timeout): + self.requests.append((url, content, dict(headers))) + self.thread_ids.append(threading.get_ident()) + return self._response + + +@pytest.fixture +def clean_anthropic_env(monkeypatch): + for name in ANTHROPIC_ENV_VARS: + monkeypatch.delenv(name, raising=False) + + +@pytest.fixture +def wif_engine(monkeypatch, clean_anthropic_env): + """Route the wiring's WIF tier through a fresh engine (never the module + singleton, to avoid cross-test cache pollution) and count its consultations.""" + import httpx + + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + poster = RecordingPoster( + httpx.Response( + 200, + json={"access_token": FAKE_MINTED_TOKEN, "token_type": "Bearer", "expires_in": 3600}, + ) + ) + engine = JwtBearerTokenExchangeEngine(poster=poster) + calls = [] + + def with_injected_engine(litellm_params, api_base, model): + calls.append(model) + return get_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "get_anthropic_wif_token", with_injected_engine) + return poster, calls + + +@pytest.fixture +def wif_async_engine(monkeypatch, clean_anthropic_env): + """Route both WIF facades through one fresh engine; the poster records the + thread each exchange ran on and sync-facade consultations are counted so + async tests can prove the mint went through the async seam, off the loop.""" + import httpx + + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import aget_anthropic_wif_token, get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + poster = RecordingPoster( + httpx.Response( + 200, + json={"access_token": FAKE_MINTED_TOKEN, "token_type": "Bearer", "expires_in": 3600}, + ) + ) + engine = JwtBearerTokenExchangeEngine(poster=poster) + sync_calls = [] + + def sync_shim(litellm_params, api_base, model): + sync_calls.append(model) + return get_anthropic_wif_token(litellm_params, api_base, model, engine) + + async def async_shim(litellm_params, api_base, model): + return await aget_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "get_anthropic_wif_token", sync_shim) + monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim) + return poster, sync_calls + + +def _validate_chat_environment(api_key=None): + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + return AnthropicModelInfo().validate_environment( + headers={}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + api_key=api_key, + api_base=None, + ) + + +class TestWifTierPrecedence: + """WIF is the LOWEST credential tier: any api_key / auth_token source must + win without the engine ever being consulted.""" + + def _set_wif_env(self, monkeypatch): + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + + def test_explicit_api_key_beats_wif(self, monkeypatch, wif_engine): + poster, calls = wif_engine + self._set_wif_env(monkeypatch) + + headers = _validate_chat_environment(api_key=FAKE_REGULAR_KEY) + + assert headers["x-api-key"] == FAKE_REGULAR_KEY + assert calls == [] + assert poster.requests == [] + + def test_api_key_env_beats_wif(self, monkeypatch, wif_engine): + poster, calls = wif_engine + self._set_wif_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + + headers = _validate_chat_environment() + + assert headers["x-api-key"] == FAKE_REGULAR_KEY + assert calls == [] + assert poster.requests == [] + + def test_auth_token_env_beats_wif(self, monkeypatch, wif_engine): + poster, calls = wif_engine + self._set_wif_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_AUTH_TOKEN", FAKE_AUTH_TOKEN) + + headers = _validate_chat_environment() + + assert headers["authorization"] == f"Bearer {FAKE_AUTH_TOKEN}" + assert calls == [] + assert poster.requests == [] + + def test_wif_alone_mints_once(self, monkeypatch, wif_engine): + poster, calls = wif_engine + self._set_wif_env(monkeypatch) + + headers = _validate_chat_environment() + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert calls == ["claude-sonnet-4-5"] + assert len(poster.requests) == 1 + assert poster.requests[0][0] == "https://api.anthropic.com/v1/oauth/token" + + @pytest.mark.parametrize("blank", ["", " "]) + def test_a_blank_api_key_env_falls_through_to_wif(self, monkeypatch, wif_engine, blank): + """An empty ANTHROPIC_API_KEY cannot authenticate anything, so treating it as set would leave a + federated deployment sending an empty x-api-key on every call instead of a minted token.""" + poster, calls = wif_engine + self._set_wif_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_API_KEY", blank) + + headers = _validate_chat_environment() + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert "x-api-key" not in headers + assert calls == ["claude-sonnet-4-5"] + + @pytest.mark.parametrize("blank", ["", " "]) + def test_a_blank_auth_token_env_falls_through_to_wif(self, monkeypatch, wif_engine, blank): + poster, calls = wif_engine + self._set_wif_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_AUTH_TOKEN", blank) + + headers = _validate_chat_environment() + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert calls == ["claude-sonnet-4-5"] + + def test_a_static_key_on_a_federated_deployment_warns_once(self, monkeypatch, wif_engine, caplog): + """Static credentials outrank federation everywhere in the provider, so an operator who + configured federation and left a key behind gets no other signal that nothing is federated.""" + import logging + + from litellm.llms.anthropic.wif import _warn_static_credential_shadows_federation + + _warn_static_credential_shadows_federation.cache_clear() + poster, calls = wif_engine + self._set_wif_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + first = _validate_chat_environment() + _validate_chat_environment() + + shadow_warnings = [record for record in caplog.records if "takes precedence" in record.getMessage()] + assert first["x-api-key"] == FAKE_REGULAR_KEY + assert calls == [] + assert len(shadow_warnings) == 1, "a per-request warning would drown the log it is meant to reach" + assert "claude-sonnet-4-5" in shadow_warnings[0].getMessage() + + def test_an_unfederated_deployment_stays_quiet(self, monkeypatch, wif_engine, caplog): + import logging + + from litellm.llms.anthropic.wif import _warn_static_credential_shadows_federation + + _warn_static_credential_shadows_federation.cache_clear() + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + _validate_chat_environment() + + assert [record for record in caplog.records if "takes precedence" in record.getMessage()] == [] + + def test_the_shadow_check_reads_its_environment_once_per_deployment(self, monkeypatch, wif_engine, caplog): + """Every static-key Anthropic call reaches the shadow check, so resolving the federation rule + id from the environment per request would put a secret-manager read, and on a miss an ERROR + with a traceback, in front of every one of them. The environment is read once per deployment + instead, which is why a rule id appearing later in the process does not start warning.""" + import logging + + from litellm.llms.anthropic.wif import _warn_static_credential_shadows_federation + + _warn_static_credential_shadows_federation.cache_clear() + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + _validate_chat_environment() + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_wire1") + _validate_chat_environment() + + assert [record for record in caplog.records if "takes precedence" in record.getMessage()] == [] + + + +class TestWifZeroBehaviorChange: + def test_unconfigured_raises_same_authentication_error(self, clean_anthropic_env): + """No WIF config and no keys: same AuthenticationError as today (message + extended, type and provider identical).""" + import litellm + + with pytest.raises(litellm.AuthenticationError) as exc_info: + _validate_chat_environment() + + assert exc_info.value.llm_provider == "anthropic" + assert "ANTHROPIC_API_KEY" in exc_info.value.message + assert "ANTHROPIC_AUTH_TOKEN" in exc_info.value.message + assert "ANTHROPIC_FEDERATION_RULE_ID" in exc_info.value.message + assert "ANTHROPIC_ORGANIZATION_ID" in exc_info.value.message + assert "ANTHROPIC_SERVICE_ACCOUNT_ID" in exc_info.value.message + assert "ANTHROPIC_IDENTITY_TOKEN_FILE" in exc_info.value.message + + +class TestWifHeaderContract: + def test_minted_token_headers(self, monkeypatch, wif_engine): + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + + headers = _validate_chat_environment() + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert "oauth-2025-04-20" in headers["anthropic-beta"] + assert "x-api-key" not in headers + assert "anthropic-dangerous-direct-browser-access" not in headers + + def test_consumer_oat_key_keeps_dangerous_header(self, clean_anthropic_env): + """Regression: user-supplied consumer sk-ant-oat keys keep today's behavior.""" + headers = _validate_chat_environment(api_key=FAKE_OAUTH_TOKEN) + + assert headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + assert headers["anthropic-dangerous-direct-browser-access"] == "true" + assert "oauth-2025-04-20" in headers["anthropic-beta"] + + +class TestMergeAnthropicBetaHeaders: + """The Skills surface accepted a list-valued anthropic-beta before it shared this helper, + so the helper has to keep taking one: .split() on a list is an AttributeError.""" + + def test_list_valued_existing_header_is_merged(self): + from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers + + assert merge_anthropic_beta_headers(["skills-2025-10-02", "files-api-2025-04-14"], "oauth-2025-04-20") == ( + "files-api-2025-04-14,oauth-2025-04-20,skills-2025-10-02" + ) + + def test_list_and_comma_string_forms_agree(self): + from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers + + as_list = merge_anthropic_beta_headers(["a", "b"], "c") + as_string = merge_anthropic_beta_headers("a,b", "c") + assert as_list == as_string == "a,b,c" + + def test_skills_validate_environment_accepts_a_list_header(self, monkeypatch): + """End of the regression: the Skills surface itself must not raise on the list form.""" + from litellm.llms.anthropic.skills.transformation import AnthropicSkillsConfig + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + + headers = AnthropicSkillsConfig().validate_environment( + headers={"anthropic-beta": ["files-api-2025-04-14"]}, + litellm_params=None, + ) + + assert "files-api-2025-04-14" in headers["anthropic-beta"] + assert isinstance(headers["anthropic-beta"], str) + + +class TestWifServerOwnedAuthHeaderStrip: + """A WIF-minted token must never ride alongside a caller-supplied credential + header, but that stripping must fire only when a mint actually happened.""" + + def test_mint_strips_caller_supplied_x_api_key(self, monkeypatch, wif_engine): + """Security regression: without the strip, a caller-forwarded x-api-key + would sit next to the server-minted Authorization on the outgoing request.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + poster, _ = wif_engine + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + caller_key = "sk-ant-CALLER-SUPPLIED" + + headers = AnthropicModelInfo().validate_environment( + headers={"x-api-key": caller_key}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert "x-api-key" not in headers + assert caller_key not in headers.values() + assert len(poster.requests) == 1 + + def test_server_owned_set_is_every_proxy_credential_header(self): + """The strip list must track the proxy's own key-header list, not a hand-rolled + pair: every header user_api_key_auth accepts a LiteLLM key in must be here.""" + from litellm.llms.anthropic.common_utils import _SERVER_OWNED_AUTH_HEADERS + + assert _SERVER_OWNED_AUTH_HEADERS == SpecialHeaders.litellm_credential_header_names() + assert {"x-litellm-api-key", "api-key", "x-goog-api-key"} < _SERVER_OWNED_AUTH_HEADERS + + @pytest.mark.parametrize("header_name", PROXY_CREDENTIAL_HEADER_NAMES) + def test_mint_strips_every_proxy_credential_header(self, monkeypatch, wif_engine, header_name): + """A LiteLLM virtual key arrives in any of the proxy's accepted key headers; once + a mint happened none of them may reach Anthropic in any header slot.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + caller_key = "sk-litellm-CALLER-VIRTUAL-KEY" + + headers = AnthropicModelInfo().validate_environment( + headers={header_name.title(): caller_key, "user-agent": "caller/1.0"}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert header_name == "authorization" or header_name not in {name.lower() for name in headers} + assert all(caller_key not in value for value in headers.values()) + assert headers["user-agent"] == "caller/1.0" + + @pytest.mark.parametrize("header_name", PROXY_CREDENTIAL_HEADER_NAMES) + def test_skills_surface_strips_caller_credentials_too(self, monkeypatch, wif_engine, header_name): + """Skills builds its own headers as well; every minting surface needs the same strip.""" + from litellm.llms.anthropic.skills.transformation import AnthropicSkillsConfig + + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + caller_key = "sk-litellm-CALLER-VIRTUAL-KEY" + + headers = AnthropicSkillsConfig().validate_environment( + headers={header_name.title(): caller_key, "user-agent": "caller/1.0"}, + litellm_params=None, + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert header_name == "authorization" or header_name not in {name.lower() for name in headers} + assert all(caller_key not in value for value in headers.values()) + assert headers["user-agent"] == "caller/1.0" + + def test_passthrough_honors_a_case_variant_caller_key_instead_of_minting(self, monkeypatch, wif_engine): + """The passthrough surface hands the caller's own credential upstream rather than minting. + That check was case-sensitive, so X-Api-Key slipped past it and the caller's key would have + travelled beside a minted Bearer.""" + from litellm.llms.anthropic.pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + poster, _ = wif_engine + caller_key = "sk-ant-CALLER-SUPPLIED" + + headers, _ = AnthropicMessagesConfig().validate_anthropic_messages_environment( + headers={"X-Api-Key": caller_key}, + model="claude-sonnet-4-5", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["X-Api-Key"] == caller_key + assert "authorization" not in {name.lower() for name in headers} + assert len(poster.requests) == 0 + + @pytest.mark.parametrize("header_name", PROXY_CREDENTIAL_HEADER_NAMES) + def test_batches_surface_strips_caller_credentials_too(self, monkeypatch, wif_engine, header_name): + """Batches builds its own headers on the create path, so it needs the same strip: the + handler's retrieve path passes none, but this entry point takes the caller's.""" + from litellm.llms.anthropic.batches.transformation import AnthropicBatchesConfig + + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + caller_key = "sk-litellm-CALLER-VIRTUAL-KEY" + + headers = AnthropicBatchesConfig().validate_environment( + headers={header_name.title(): caller_key, "user-agent": "caller/1.0"}, + model="claude-sonnet-4-5", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert header_name == "authorization" or header_name not in {name.lower() for name in headers} + assert all(caller_key not in value for value in headers.values()) + assert headers["user-agent"] == "caller/1.0" + + @pytest.mark.parametrize("header_name", PROXY_CREDENTIAL_HEADER_NAMES) + def test_files_surface_strips_caller_credentials_too(self, monkeypatch, wif_engine, header_name): + """The files surface builds its own headers, so it needs the same strip the chat surface + has: without it a minted federation Bearer travels beside the caller's own credential.""" + from litellm.llms.anthropic.files.transformation import AnthropicFilesConfig + + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + caller_key = "sk-litellm-CALLER-VIRTUAL-KEY" + + headers = AnthropicFilesConfig().validate_environment( + headers={header_name.title(): caller_key, "user-agent": "caller/1.0"}, + model="claude-sonnet-4-5", + messages=[], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert header_name == "authorization" or header_name not in {name.lower() for name in headers} + assert all(caller_key not in value for value in headers.values()) + assert headers["user-agent"] == "caller/1.0" + + def test_no_mint_preserves_caller_supplied_authorization(self, monkeypatch, clean_anthropic_env): + """No-regression: LiteLLM deliberately lets a caller-forwarded credential + header ride alongside a statically configured ANTHROPIC_API_KEY, because the + two occupy different header slots (x-api-key vs authorization) when the key + isn't OAuth-shaped. The strip must stay conditional on an actual WIF mint.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + caller_authorization = "Bearer caller-forwarded-downstream-token" + + headers = AnthropicModelInfo().validate_environment( + headers={"authorization": caller_authorization}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["x-api-key"] == FAKE_REGULAR_KEY + assert headers["authorization"] == caller_authorization + + +class TestWifResolvedApiKeyThreading: + """Regression for the resolved_api_key local (formerly a rebind of the api_key + parameter): a minted token must reach the outgoing headers on both the sync + validate_environment path and the async aget_auth_header path, never a stale + None left over from the original unresolved parameter.""" + + def test_validate_environment_carries_minted_token(self, monkeypatch, wif_engine): + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + + headers = _validate_chat_environment() + + assert "authorization" in headers + assert headers["authorization"] not in (None, "Bearer None") + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + + @pytest.mark.asyncio + async def test_aget_auth_header_carries_minted_token(self, monkeypatch, wif_async_engine): + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + + result = await AnthropicModelInfo.aget_auth_header(allow_workload_identity=True) + + assert result is not None + assert result["authorization"] not in (None, "Bearer None") + assert result["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + + +class TestGetAuthHeaderBetas: + def test_oat_branch_carries_oauth_beta(self, clean_anthropic_env): + """The pre-existing bug: the oat branch returned a bare Bearer without the + mandatory oauth beta.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + result = AnthropicModelInfo.get_auth_header(api_key=FAKE_OAUTH_TOKEN) + + assert result == { + "authorization": f"Bearer {FAKE_OAUTH_TOKEN}", + "anthropic-beta": "oauth-2025-04-20", + } + + def test_wif_fallback_returns_bearer_and_beta(self, monkeypatch, wif_engine): + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + poster, calls = wif_engine + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + + result = AnthropicModelInfo.get_auth_header(allow_workload_identity=True) + + assert result == { + "authorization": f"Bearer {FAKE_MINTED_TOKEN}", + "anthropic-beta": "oauth-2025-04-20", + } + assert len(poster.requests) == 1 + + def test_no_credentials_still_returns_none(self, clean_anthropic_env): + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + assert AnthropicModelInfo.get_auth_header() is None + + +class TestFilesBatchesBetaMerge: + """Regression for the anthropic-beta clobber (files) and drop (batches): + a Bearer oat auth header must keep the oauth beta AND gain the surface beta.""" + + def test_files_merges_oauth_and_files_betas(self, clean_anthropic_env): + from litellm.llms.anthropic.files.transformation import AnthropicFilesConfig + + headers = AnthropicFilesConfig().validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key=FAKE_OAUTH_TOKEN, + ) + + assert headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + betas = set(headers["anthropic-beta"].split(",")) + assert {"oauth-2025-04-20", "files-api-2025-04-14"} <= betas + + def test_batches_merges_oauth_and_batches_betas(self, clean_anthropic_env): + from litellm.llms.anthropic.batches.transformation import AnthropicBatchesConfig + + headers = AnthropicBatchesConfig().validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key=FAKE_OAUTH_TOKEN, + ) + + assert headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + betas = set(headers["anthropic-beta"].split(",")) + assert {"oauth-2025-04-20", "message-batches-2024-09-24"} <= betas + + def test_files_preserves_caller_supplied_beta(self, clean_anthropic_env): + """Regression: files did a two-way merge that dropped the client's own + anthropic-beta; it must three-way merge exactly like batches.""" + from litellm.llms.anthropic.files.transformation import AnthropicFilesConfig + + headers = AnthropicFilesConfig().validate_environment( + headers={"anthropic-beta": "context-1m-2025-08-07"}, + model="", + messages=[], + optional_params={}, + litellm_params={}, + api_key=FAKE_OAUTH_TOKEN, + ) + + betas = set(headers["anthropic-beta"].split(",")) + assert {"context-1m-2025-08-07", "oauth-2025-04-20", "files-api-2025-04-14"} <= betas + + +class TestMessagesEnvAuthBetaMerge: + def test_client_beta_survives_env_auth_injection(self, monkeypatch, clean_anthropic_env): + """Regression: headers.update(auth_header) silently clobbered the client's + anthropic-beta on the native /v1/messages route.""" + from litellm.llms.anthropic.pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_OAUTH_TOKEN) + + headers, _ = AnthropicMessagesConfig().validate_anthropic_messages_environment( + headers={"anthropic-beta": "context-1m-2025-08-07"}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={}, + ) + + assert headers["authorization"] == f"Bearer {FAKE_OAUTH_TOKEN}" + betas = set(headers["anthropic-beta"].split(",")) + assert {"context-1m-2025-08-07", "oauth-2025-04-20"} <= betas + + +WIF_PARAMS_ONLY = { + "anthropic_federation_rule_id": "fdrl_params", + "anthropic_organization_id": "org-params", + "anthropic_identity_token": "oidc/env/WIF_PARAMS_TEST_TOKEN", +} + + +class TestWifLitellmParamsPlumbing: + """Per-deployment anthropic_* litellm_params must reach the WIF tier on every + surface that has them, not only chat.""" + + @pytest.fixture(autouse=True) + def _inline_identity_token(self, monkeypatch): + monkeypatch.setenv("WIF_PARAMS_TEST_TOKEN", "params-jwt") + + def test_files_mints_from_litellm_params(self, wif_engine): + from litellm.llms.anthropic.files.transformation import AnthropicFilesConfig + + poster, _ = wif_engine + headers = AnthropicFilesConfig().validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params=dict(WIF_PARAMS_ONLY), + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert len(poster.requests) == 1 + + def test_batches_mints_from_litellm_params(self, wif_engine): + from litellm.llms.anthropic.batches.transformation import AnthropicBatchesConfig + + poster, _ = wif_engine + headers = AnthropicBatchesConfig().validate_environment( + headers={}, + model="", + messages=[], + optional_params={}, + litellm_params=dict(WIF_PARAMS_ONLY), + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert len(poster.requests) == 1 + + def test_skills_mints_from_litellm_params(self, wif_engine): + from litellm.llms.anthropic.skills.transformation import AnthropicSkillsConfig + from litellm.types.router import GenericLiteLLMParams + + poster, _ = wif_engine + headers = AnthropicSkillsConfig().validate_environment( + headers={}, + litellm_params=GenericLiteLLMParams( + anthropic_federation_rule_id=WIF_PARAMS_ONLY["anthropic_federation_rule_id"], + anthropic_organization_id=WIF_PARAMS_ONLY["anthropic_organization_id"], + anthropic_identity_token=WIF_PARAMS_ONLY["anthropic_identity_token"], + ), + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert len(poster.requests) == 1 + + def test_messages_mints_from_litellm_params(self, wif_engine): + from litellm.llms.anthropic.pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + poster, _ = wif_engine + headers, _ = AnthropicMessagesConfig().validate_anthropic_messages_environment( + headers={}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params=dict(WIF_PARAMS_ONLY), + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert len(poster.requests) == 1 + + +class TestWifTokenUrlParity: + """Both credential tiers must derive the SAME clean token URL from any form of + the deployment base; a mismatch also duplicates mints because token_url is in + the engine cache key.""" + + @pytest.mark.parametrize( + "configured_base", + [ + "https://gw.example.com", + "https://gw.example.com/", + "https://gw.example.com/v1/messages", + "https://gw.example.com/v1/messages/", + ], + ) + def test_both_tiers_share_one_clean_token_url(self, monkeypatch, wif_engine, configured_base): + # This is about deriving one URL from many spellings of the same base, not about which + # hosts an operator trusts with org-scoped credentials, so the private host is allowlisted. + monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gw.example.com") + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + poster, _ = wif_engine + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + + AnthropicModelInfo().validate_environment( + headers={}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params={"api_base": configured_base}, + api_key=None, + api_base=None, + ) + AnthropicModelInfo.get_auth_header(api_base=configured_base, allow_workload_identity=True) + + assert [url for (url, _, _) in poster.requests] == ["https://gw.example.com/v1/oauth/token"] + + +class TestWifAsyncSeam: + """Async callers must resolve the WIF tier through the async facade so a cold + mint never blocks the event loop.""" + + @pytest.mark.asyncio + async def test_aget_auth_header_runs_exchange_off_event_loop(self, monkeypatch, wif_async_engine): + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + poster, sync_calls = wif_async_engine + for name, value in WIF_ENV.items(): + monkeypatch.setenv(name, value) + + result = await AnthropicModelInfo.aget_auth_header(allow_workload_identity=True) + + assert result == { + "authorization": f"Bearer {FAKE_MINTED_TOKEN}", + "anthropic-beta": "oauth-2025-04-20", + } + assert sync_calls == [] + assert poster.thread_ids == [poster.thread_ids[0]] + assert poster.thread_ids[0] != threading.get_ident() + + @pytest.mark.asyncio + async def test_avalidate_messages_environment_mints_off_loop(self, wif_async_engine, monkeypatch): + from litellm.llms.anthropic.pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + poster, sync_calls = wif_async_engine + monkeypatch.setenv("WIF_PARAMS_TEST_TOKEN", "params-jwt") + + headers, _ = await AnthropicMessagesConfig().avalidate_anthropic_messages_environment( + headers={}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, + litellm_params=dict(WIF_PARAMS_ONLY), + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert sync_calls == [] + assert poster.thread_ids[0] != threading.get_ident() + + @pytest.mark.asyncio + async def test_avalidate_delegates_to_subclass_sync_override(self): + """A provider subclass that only overrides the sync method must keep its + behavior when the handler goes through the async variant.""" + from litellm.llms.anthropic.pass_through.messages.transformation import ( + AnthropicMessagesConfig, + ) + + class MarkerConfig(AnthropicMessagesConfig): + def validate_anthropic_messages_environment( + self, + headers, + model, + messages, + optional_params, + litellm_params, + api_key=None, + api_base=None, + ): + return {"x-marker": "sync"}, api_base + + headers, api_base = await MarkerConfig().avalidate_anthropic_messages_environment( + headers={}, + model="claude-sonnet-4-5", + messages=[], + optional_params={}, + litellm_params={}, + api_base="https://marker.example.com", + ) + + assert headers == {"x-marker": "sync"} + assert api_base == "https://marker.example.com" + + @pytest.mark.asyncio + async def test_base_default_avalidate_delegates_to_sync(self): + from litellm.llms.base_llm.anthropic_messages.transformation import ( + BaseAnthropicMessagesConfig, + ) + + class SyncOnlyConfig(BaseAnthropicMessagesConfig): + def validate_anthropic_messages_environment( + self, + headers, + model, + messages, + optional_params, + litellm_params, + api_key=None, + api_base=None, + ): + return {"x-sync-only": "1"}, api_base + + def get_complete_url(self, api_base, api_key, model, optional_params, litellm_params, stream=None): + return api_base or "" + + def get_supported_anthropic_messages_params(self, model): + return [] + + def transform_anthropic_messages_request( + self, model, messages, anthropic_messages_optional_request_params, litellm_params, headers + ): + return {} + + def transform_anthropic_messages_response(self, model, raw_response, logging_obj): + raise NotImplementedError + + headers, _ = await SyncOnlyConfig().avalidate_anthropic_messages_environment( + headers={}, + model="m", + messages=[], + optional_params={}, + litellm_params={}, + ) + + assert headers == {"x-sync-only": "1"} + + +class TestWifRespxEndToEnd: + def test_completion_mints_and_never_leaks_config(self, monkeypatch, tmp_path, clean_anthropic_env): + """Drives the REAL kwargs funnel through litellm.completion: the mint hits + /v1/oauth/token, the data plane carries the minted Bearer + oauth beta, and + NONE of the six anthropic_* keys leak into the /v1/messages body.""" + import httpx + import respx + + import litellm + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.setattr(litellm, "anthropic_key", None) + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = tmp_path / "identity-token" + token_file.write_text("e2e-oidc-assertion", encoding="utf-8") + + engine = JwtBearerTokenExchangeEngine() + monkeypatch.setattr( + anthropic_common_utils, + "get_anthropic_wif_token", + lambda litellm_params, api_base, model: get_anthropic_wif_token(litellm_params, api_base, model, engine), + ) + + wif_kwarg_names: Final = ( + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_service_account_id", + "anthropic_federation_workspace_id", + "anthropic_identity_token_file", + "anthropic_identity_token", + ) + anthropic_response = { + "id": "msg_01", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "Hello from WIF"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 5, "output_tokens": 4}, + } + + with respx.mock: + token_route = respx.post("https://api.anthropic.com/v1/oauth/token").mock( + return_value=httpx.Response( + 200, + json={"access_token": FAKE_MINTED_TOKEN, "token_type": "Bearer", "expires_in": 3600}, + ) + ) + messages_route = respx.post("https://api.anthropic.com/v1/messages").mock( + return_value=httpx.Response(200, json=anthropic_response) + ) + response = litellm.completion( + model="anthropic/claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + anthropic_federation_rule_id="fdrl_e2e", + anthropic_organization_id="org-e2e", + anthropic_service_account_id="svcacct_e2e", + anthropic_federation_workspace_id="wrkspc_e2e", + anthropic_identity_token_file=str(token_file), + anthropic_identity_token="oidc/env/UNUSED_FALLBACK", + ) + + assert response.choices[0].message.content == "Hello from WIF" + assert token_route.call_count == 1 + exchange_body = json.loads(token_route.calls[0].request.content) + assert exchange_body["assertion"] == "e2e-oidc-assertion" + assert exchange_body["federation_rule_id"] == "fdrl_e2e" + + data_request = messages_route.calls[0].request + assert data_request.headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert "oauth-2025-04-20" in data_request.headers["anthropic-beta"] + assert "x-api-key" not in data_request.headers + assert "anthropic-dangerous-direct-browser-access" not in data_request.headers + data_body = json.loads(data_request.content) + for key in wif_kwarg_names: + assert key not in data_body + + def test_completion_with_trailing_slash_api_base_mints_at_clean_token_url( + self, monkeypatch, tmp_path, clean_anthropic_env + ): + """Regression: a trailing-slash api_base defeated the endswith check in + main.py AND the removesuffix surgery, sending the exchange POST to + .../v1/messages/v1/oauth/token (404).""" + import httpx + import respx + + import litellm + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + monkeypatch.setattr(litellm, "api_key", None) + monkeypatch.setattr(litellm, "anthropic_key", None) + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = tmp_path / "identity-token" + token_file.write_text("e2e-oidc-assertion", encoding="utf-8") + + engine = JwtBearerTokenExchangeEngine() + monkeypatch.setattr( + anthropic_common_utils, + "get_anthropic_wif_token", + lambda litellm_params, api_base, model: get_anthropic_wif_token(litellm_params, api_base, model, engine), + ) + + anthropic_response = { + "id": "msg_02", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "Hello from WIF"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 5, "output_tokens": 4}, + } + + with respx.mock: + token_route = respx.post("https://api.anthropic.com/v1/oauth/token").mock( + return_value=httpx.Response( + 200, + json={"access_token": FAKE_MINTED_TOKEN, "token_type": "Bearer", "expires_in": 3600}, + ) + ) + messages_route = respx.post(url__regex=r"https://api\.anthropic\.com/v1/messages.*").mock( + return_value=httpx.Response(200, json=anthropic_response) + ) + response = litellm.completion( + model="anthropic/claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + api_base="https://api.anthropic.com/v1/messages/", + anthropic_federation_rule_id="fdrl_e2e", + anthropic_organization_id="org-e2e", + anthropic_identity_token_file=str(token_file), + ) + + assert response.choices[0].message.content == "Hello from WIF" + assert token_route.call_count == 1 + assert messages_route.calls[0].request.headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + + def test_get_auth_header_with_litellm_params_mints_via_real_engine( + self, monkeypatch, tmp_path, clean_anthropic_env + ): + import httpx + import respx + + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + from litellm.llms.anthropic.wif import get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = tmp_path / "identity-token" + token_file.write_text("e2e-oidc-assertion", encoding="utf-8") + + engine = JwtBearerTokenExchangeEngine() + monkeypatch.setattr( + anthropic_common_utils, + "get_anthropic_wif_token", + lambda litellm_params, api_base, model: get_anthropic_wif_token(litellm_params, api_base, model, engine), + ) + + with respx.mock: + token_route = respx.post("https://api.anthropic.com/v1/oauth/token").mock( + return_value=httpx.Response( + 200, + json={"access_token": FAKE_MINTED_TOKEN, "token_type": "Bearer", "expires_in": 3600}, + ) + ) + result = AnthropicModelInfo.get_auth_header( + allow_workload_identity=True, + litellm_params={ + "anthropic_federation_rule_id": "fdrl_e2e", + "anthropic_organization_id": "org-e2e", + "anthropic_identity_token_file": str(token_file), + }, + ) + + assert result == { + "authorization": f"Bearer {FAKE_MINTED_TOKEN}", + "anthropic-beta": "oauth-2025-04-20", + } + assert token_route.call_count == 1 + exchange_body = json.loads(token_route.calls[0].request.content) + assert exchange_body["federation_rule_id"] == "fdrl_e2e" + + +class TestWifProviderAllowlist: + """A federation token is an Anthropic-org credential, and the exchange POSTs the workload's OIDC + assertion to the deployment's own api_base host. Providers that subclass the Anthropic config for + their own endpoints must therefore never reach the WIF tier, even when it is configured purely + through ANTHROPIC_* environment variables.""" + + @staticmethod + def _env_only_wif(monkeypatch) -> None: # noqa: D401 + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_prod") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-prod-uuid") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "oidc/env/WIF_TEST_JWT") + monkeypatch.setenv("WIF_TEST_JWT", "jwt-assertion-value") + + def test_vertex_anthropic_never_mints_or_sends_the_assertion(self, monkeypatch, wif_engine): + from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation import ( + VertexAIAnthropicConfig, + ) + + import litellm + + poster, calls = wif_engine + self._env_only_wif(monkeypatch) + + with pytest.raises(litellm.AuthenticationError): + VertexAIAnthropicConfig().validate_environment( + headers={}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + api_key=None, + api_base="https://us-east5-aiplatform.googleapis.com/v1/projects/p/locations/us-east5", + ) + + assert calls == [] + assert poster.requests == [] + + def test_anthropic_itself_still_mints(self, monkeypatch, wif_engine): + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + poster, calls = wif_engine + self._env_only_wif(monkeypatch) + + headers = AnthropicConfig().validate_environment( + headers={}, + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + optional_params={}, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + assert len(poster.requests) == 1 + + def test_auth_header_facade_defaults_to_refusing_to_mint(self, monkeypatch, clean_anthropic_env): + """The facade is reachable from provider code that has nothing to do with Anthropic, so a + caller must state that it authenticates against Anthropic's own API.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + self._env_only_wif(monkeypatch) + + assert AnthropicModelInfo.get_auth_header(None) is None + assert AnthropicModelInfo.get_auth_header(None, allow_workload_identity=False) is None + + def test_eligibility_is_not_inherited_by_a_new_subclass(self): + """A provider added later by subclassing the Anthropic config must not inherit the right to + mint an Anthropic-org credential against its own host.""" + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + from litellm.llms.anthropic.common_utils import config_allows_workload_identity + + class NewCompatibleProvider(AnthropicConfig): + pass + + assert config_allows_workload_identity(AnthropicConfig()) is True + assert config_allows_workload_identity(NewCompatibleProvider()) is False + + def test_model_discovery_gates_on_the_instance(self, monkeypatch, wif_engine): + """get_models is inherited, so it must consult the instance rather than trusting its caller.""" + from litellm.llms.vertex_ai.vertex_ai_partner_models.anthropic.transformation import ( + VertexAIAnthropicConfig, + ) + + poster, calls = wif_engine + self._env_only_wif(monkeypatch) + + with pytest.raises(ValueError, match="ANTHROPIC_API_KEY"): + VertexAIAnthropicConfig().get_models( + api_base="https://us-east5-aiplatform.googleapis.com/v1/projects/p/locations/us-east5" + ) + + assert poster.requests == [] + + +def _models_page_response(page: dict, status_code: int = 200): + import httpx + + return httpx.Response(status_code, json=page, request=httpx.Request("GET", "https://api.anthropic.com/v1/models")) + + +class RecordingModelsClient: + """Records every call and answers with the queued responses in order, cycling the last one + once exhausted so a runaway pagination loop degrades to a repeated page rather than an + IndexError, letting the page-cap test observe the cap firing instead of a test bug.""" + + def __init__(self, pages: list[dict] | None = None, responses=None): + self.calls = [] + self._responses = responses if responses is not None else [_models_page_response(page) for page in pages] + + def get(self, url, headers=None, params=None, follow_redirects=None, timeout=None): + self.calls.append(SimpleNamespace(url=url, headers=headers, params=params, follow_redirects=follow_redirects)) + index = min(len(self.calls) - 1, len(self._responses) - 1) + return self._responses[index] + + +class TestModelDiscovery: + """AnthropicModelInfo.get_models / discover_models: pagination, redirect refusal, and + sanitized errors on the upstream Anthropic /v1/models call itself (issue #28607 gap: a + WIF source configured in litellm_params, rather than the environment, could not + discover).""" + + @pytest.mark.parametrize( + "configured_base", ["https://api.anthropic.com/v1", "https://api.anthropic.com/v1/messages"] + ) + def test_discovery_does_not_double_the_version_segment(self, monkeypatch, clean_anthropic_env, configured_base): + """Regression: /v1/models is appended here, so a base an operator already wrote as + .../v1 (or the chat URL they copied) would be asked for /v1/v1/models and 404.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient([{"data": [{"id": "claude-a"}], "has_more": False, "last_id": "claude-a"}]) + monkeypatch.setattr("litellm.module_level_client", client) + + models = AnthropicModelInfo().get_models(api_base=configured_base) + + assert models == ["anthropic/claude-a"] + assert client.calls[0].url == "https://api.anthropic.com/v1/models" + + def test_get_models_paginates_via_has_more_and_last_id(self, monkeypatch, clean_anthropic_env): + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient( + [ + {"data": [{"id": "claude-a"}, {"id": "claude-b"}], "has_more": True, "last_id": "claude-b"}, + {"data": [{"id": "claude-c"}], "has_more": False, "last_id": "claude-c"}, + ] + ) + monkeypatch.setattr("litellm.module_level_client", client) + + models = AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert models == ["anthropic/claude-a", "anthropic/claude-b", "anthropic/claude-c"] + assert len(client.calls) == 2 + assert client.calls[0].url == "https://api.anthropic.com/v1/models" + assert client.calls[1].url == "https://api.anthropic.com/v1/models?after_id=claude-b" + + def test_paginated_fetch_survives_the_real_http_client(self, monkeypatch, clean_anthropic_env): + """Regression: the second page is fetched through the real HTTPHandler, which merges the + URL's query string into the params mapping by mutating it. Handing that client a + read-only mapping raised AttributeError, so discovery blew up for any org holding more + models than one page, while the stubbed client here never exercised the mutation.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + from litellm.llms.custom_httpx.http_handler import HTTPHandler + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + requested: Final = [] # mutable-ok: a test spy recording the URLs the client was asked for + + def respond(request: httpx.Request) -> httpx.Response: + requested.append(str(request.url)) + first: Final = "after_id" not in request.url.params + return httpx.Response( + 200, + json={ + "data": [{"id": "claude-a"}] if first else [{"id": "claude-b"}], + "has_more": first, + "last_id": "claude-a" if first else "claude-b", + }, + ) + + handler: Final = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(respond))) + monkeypatch.setattr("litellm.module_level_client", handler) + + models = AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert models == ["anthropic/claude-a", "anthropic/claude-b"] + assert requested == [ + "https://api.anthropic.com/v1/models", + "https://api.anthropic.com/v1/models?after_id=claude-a", + ] + + def test_get_models_refuses_to_follow_redirects(self, monkeypatch, clean_anthropic_env): + """Only the configured api_base is validated, so a redirected /v1/models must not be + allowed to replay the credential to an unvalidated origin -- same rule already applied + to the WIF token exchange itself.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient([{"data": [], "has_more": False, "last_id": None}]) + monkeypatch.setattr("litellm.module_level_client", client) + + AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert client.calls[0].follow_redirects is False + + def test_get_models_page_cap_stops_a_runaway_has_more(self, monkeypatch, clean_anthropic_env): + from litellm.llms.anthropic.common_utils import ( + _MODEL_LIST_PAGE_CAP, + AnthropicModelInfo, + ) + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient([{"data": [{"id": "claude-loop"}], "has_more": True, "last_id": "claude-loop"}]) + monkeypatch.setattr("litellm.module_level_client", client) + + with pytest.raises(Exception, match="did not terminate"): + AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert len(client.calls) == _MODEL_LIST_PAGE_CAP + + def test_get_models_error_is_sanitized_not_raw_response_text(self, monkeypatch, clean_anthropic_env): + """A failed discovery call must never echo the raw response body verbatim -- only the + structured error message, so an unrelated/oversized/reflected body is not surfaced.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + reflected_payload = "" * 50 + client = RecordingModelsClient( + responses=[ + _models_page_response( + { + "type": "error", + "error": { + "type": "authentication_error", + "message": "invalid x-api-key", + "reflected": reflected_payload, + }, + }, + status_code=401, + ) + ] + ) + monkeypatch.setattr("litellm.module_level_client", client) + + with pytest.raises(Exception, match="invalid x-api-key") as exc_info: # noqa: B017, PT011 # the callee raises a bare Exception; match pins the sanitized text + AnthropicModelInfo().get_models(api_base="https://api.anthropic.com") + + assert "invalid x-api-key" in str(exc_info.value) + assert reflected_payload not in str(exc_info.value) + + def test_discover_models_threads_litellm_params_into_wif(self, monkeypatch, wif_engine): + """The gap this phase fixes: get_models only ever saw api_key/api_base, so a WIF source + configured in litellm_params (rather than ANTHROPIC_* env vars) could not discover.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + poster, calls = wif_engine + client = RecordingModelsClient([{"data": [{"id": "claude-wif"}], "has_more": False, "last_id": None}]) + monkeypatch.setattr("litellm.module_level_client", client) + monkeypatch.setenv("DISC_JWT", "jwt-assertion-value") + + models = AnthropicModelInfo().discover_models( + litellm_params={ + "anthropic_federation_rule_id": "fdrl_disc", + "anthropic_organization_id": "org-disc", + "anthropic_identity_token": "oidc/env/DISC_JWT", + } + ) + + assert models == ["anthropic/claude-wif"] + assert len(poster.requests) == 1 + assert client.calls[0].headers["authorization"] == f"Bearer {FAKE_MINTED_TOKEN}" + + def test_discover_models_without_litellm_params_behaves_like_get_models(self, monkeypatch, clean_anthropic_env): + """No litellm_params (the wildcard-discovery call shape) must fall back to the + env-only resolution get_models has always used -- zero behavior change for that path.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + monkeypatch.setenv("ANTHROPIC_API_KEY", FAKE_REGULAR_KEY) + client = RecordingModelsClient([{"data": [{"id": "claude-env"}], "has_more": False, "last_id": None}]) + monkeypatch.setattr("litellm.module_level_client", client) + + models = AnthropicModelInfo().discover_models(litellm_params=None) + + assert models == ["anthropic/claude-env"] + assert client.calls[0].headers["x-api-key"] == FAKE_REGULAR_KEY + + def test_discover_models_explicit_api_key_beats_wif(self, monkeypatch, wif_engine): + """Same precedence discover_models must honor as every other Anthropic auth surface: + WIF is the lowest tier.""" + from litellm.llms.anthropic.common_utils import AnthropicModelInfo + + poster, calls = wif_engine + client = RecordingModelsClient([{"data": [], "has_more": False, "last_id": None}]) + monkeypatch.setattr("litellm.module_level_client", client) + + AnthropicModelInfo().discover_models( + litellm_params={ + "api_key": FAKE_REGULAR_KEY, + "anthropic_federation_rule_id": "fdrl_disc", + "anthropic_organization_id": "org-disc", + "anthropic_identity_token": "oidc/env/DISC_JWT", + } + ) + + assert client.calls[0].headers["x-api-key"] == FAKE_REGULAR_KEY + assert calls == [] + assert poster.requests == [] + + +class TestWifExchangeTransportHardening: + def test_token_exchange_client_does_not_follow_redirects(self): + """Only the initial token URL is validated, so a 3xx must not be allowed to replay the + assertion to an origin that was never checked.""" + from litellm.llms.base_llm.auth.token_exchange import _HttpxSyncTokenPoster + + handler = _HttpxSyncTokenPoster()._handler_instance() + + assert handler.client.follow_redirects is False + + +class TestWifServerOwnedParamsAreUnconditional: + """The minting fields choose which server-side secret is read and, with api_base, where it goes, + so no client-side credential opt-in may re-enable them.""" + + @staticmethod + def _body(param: str) -> dict: + return {"model": "claude-sonnet-5", param: "oidc/env/SOME_SERVER_SECRET"} + + @pytest.mark.parametrize( + "param", + [ + "anthropic_identity_token", + "anthropic_identity_token_file", + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_service_account_id", + # Phase 1 identity-source selection and its two variants' fields: each one + # selects a server-side secret or a destination (a signing key, a client + # secret, a token endpoint), so every one joins the same unconditional ban. + "anthropic_identity_source", + "anthropic_issuer_url", + "anthropic_issuer_subject", + "anthropic_issuer_audience", + "anthropic_issuer_ttl_seconds", + "anthropic_issuer_signing_key_ref", + "anthropic_keycloak_token_url", + "anthropic_keycloak_client_id", + "anthropic_keycloak_auth_method", + "anthropic_keycloak_client_secret_ref", + "anthropic_keycloak_scope", + ], + ) + def test_rejected_even_with_proxy_wide_opt_in(self, param: str): + from litellm.proxy.auth.auth_utils import is_request_body_safe + + with pytest.raises(ValueError, match="server-owned workload identity federation"): + is_request_body_safe( + request_body=self._body(param), + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="claude-sonnet-5", + ) + + def test_rejected_inside_nested_litellm_params(self): + from litellm.proxy.auth.auth_utils import is_request_body_safe + + with pytest.raises(ValueError, match="server-owned workload identity federation"): + is_request_body_safe( + request_body={"model": "claude-sonnet-5", "litellm_params": self._body("anthropic_identity_token")}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="claude-sonnet-5", + ) + + def test_workspace_id_is_refused_from_a_request_body(self): + """Regression, proven live against Anthropic before this was closed: a caller-supplied + workspace id reached the token endpoint, which answered "workspace_id is not a well-formed + wrkspc_ tagged ID", i.e. the caller's value had become the scope of the minted credential. + router.py merges request kwargs OVER deployment params, so it also beat the configured one.""" + from litellm.proxy.auth.auth_utils import is_request_body_safe + + with pytest.raises(Exception, match="server-owned workload identity federation parameter"): + is_request_body_safe( + request_body={"model": "claude-sonnet-5", "anthropic_federation_workspace_id": "wrkspc_abc"}, + general_settings={}, + llm_router=None, + model="claude-sonnet-5", + ) + + def test_bedrock_workspace_spellings_are_untouched(self): + """The Bedrock Claude Platform route reads its per-request workspace from these three + spellings, anthropic_workspace_id included, none of which is a federation parameter; the + federation field carries its own name, so the client-side credential opt-in that admits + them is not overridden by the unconditional federation ban.""" + from litellm.proxy.auth.auth_utils import is_request_body_safe + + for spelling in ("workspace_id", "aws_workspace_id", "anthropic_workspace_id"): + assert ( + is_request_body_safe( + request_body={"model": "claude-sonnet-5", spelling: "wrkspc_abc"}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="claude-sonnet-5", + ) + is True + ) + + +class TestWifDisabledOnClientRedirectedBase: + def test_the_sentinel_survives_the_kwargs_funnel(self): + """Setting the sentinel is only half of it. get_litellm_params rebuilds litellm_params from + kwargs, so a field it does not carry is dropped on the way and the deployment federates for + the caller-chosen base after all.""" + from litellm.litellm_core_utils.get_litellm_params import get_litellm_params + from litellm.router_utils.clientside_credential_handler import ( + DISABLE_WORKLOAD_IDENTITY_PARAM, + ) + + funneled = get_litellm_params(**{DISABLE_WORKLOAD_IDENTITY_PARAM: True}) + + assert funneled[DISABLE_WORKLOAD_IDENTITY_PARAM] is True + + def test_the_sentinel_is_not_client_settable(self): + """It is server-owned in both directions: a caller must not be able to set it, and must not + be able to clear it either.""" + from litellm.router_utils.clientside_credential_handler import ( + DISABLE_WORKLOAD_IDENTITY_PARAM, + ) + from litellm.types.router import reject_server_owned_wif_params + + with pytest.raises(ValueError, match=DISABLE_WORKLOAD_IDENTITY_PARAM): + reject_server_owned_wif_params({DISABLE_WORKLOAD_IDENTITY_PARAM: False}) + + def test_base_override_clears_wif_and_sets_the_sentinel(self): + """A federation token minted for a client-chosen api_base would send the workload's assertion, + and then the minted bearer, to that host.""" + from litellm.llms.anthropic.wif import resolve_anthropic_wif_params + from litellm.router_utils.clientside_credential_handler import ( + DISABLE_WORKLOAD_IDENTITY_PARAM, + get_dynamic_litellm_params, + ) + + admin_deployment = { + "model": "anthropic/claude-sonnet-5", + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + "anthropic_identity_token": "oidc/env/WIF_TEST_JWT", + } + + redirected = get_dynamic_litellm_params( + litellm_params=dict(admin_deployment), + request_kwargs={"api_base": "https://not-anthropic.example"}, + ) + + assert redirected[DISABLE_WORKLOAD_IDENTITY_PARAM] is True + assert "anthropic_federation_rule_id" not in redirected + assert resolve_anthropic_wif_params(redirected) is None + + def test_sentinel_blocks_env_var_configured_federation(self, monkeypatch): + """Environment-configured federation cannot be cleared out of a dict, so the sentinel is what + stops it on a redirected deployment.""" + from litellm.llms.anthropic.wif import resolve_anthropic_wif_params + from litellm.router_utils.clientside_credential_handler import DISABLE_WORKLOAD_IDENTITY_PARAM + + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "oidc/env/WIF_TEST_JWT") + monkeypatch.setenv("WIF_TEST_JWT", "jwt-assertion-value") + + assert resolve_anthropic_wif_params({}) is not None + assert resolve_anthropic_wif_params({DISABLE_WORKLOAD_IDENTITY_PARAM: True}) is None + + def test_base_override_clears_internal_issuer_fields(self): + """Same failure mode the legacy-path test above guards against, for the internal_issuer + identity source: a signing_key_ref resolved for a client-chosen api_base would mint an + assertion, and then a bearer token, for that host.""" + from litellm.llms.anthropic.wif import resolve_anthropic_wif_params + from litellm.router_utils.clientside_credential_handler import ( + DISABLE_WORKLOAD_IDENTITY_PARAM, + get_dynamic_litellm_params, + ) + + admin_deployment = { + "model": "anthropic/claude-sonnet-5", + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": "oidc/env/ISSUER_SIGNING_KEY_PEM", + } + + redirected = get_dynamic_litellm_params( + litellm_params=dict(admin_deployment), + request_kwargs={"api_base": "https://not-anthropic.example"}, + ) + + assert redirected[DISABLE_WORKLOAD_IDENTITY_PARAM] is True + assert "anthropic_identity_source" not in redirected + assert "anthropic_issuer_signing_key_ref" not in redirected + assert resolve_anthropic_wif_params(redirected) is None + + def test_base_override_clears_keycloak_fields(self): + """Same as the internal_issuer case above, for the keycloak identity source: a + client_secret_ref resolved for a client-chosen api_base must not follow it there.""" + from litellm.llms.anthropic.wif import resolve_anthropic_wif_params + from litellm.router_utils.clientside_credential_handler import ( + DISABLE_WORKLOAD_IDENTITY_PARAM, + get_dynamic_litellm_params, + ) + + admin_deployment = { + "model": "anthropic/claude-sonnet-5", + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": "https://keycloak.internal.example/realms/r/protocol/openid-connect/token", + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_client_secret_ref": "oidc/env/KEYCLOAK_CLIENT_SECRET", + } + + redirected = get_dynamic_litellm_params( + litellm_params=dict(admin_deployment), + request_kwargs={"api_base": "https://not-anthropic.example"}, + ) + + assert redirected[DISABLE_WORKLOAD_IDENTITY_PARAM] is True + assert "anthropic_identity_source" not in redirected + assert "anthropic_keycloak_client_secret_ref" not in redirected + assert resolve_anthropic_wif_params(redirected) is None + + def test_create_anthropic_model_list_response_lists_ids_as_told(): """listed_ids renames an entry for the caller while display_name and every other field stay keyed to the served id, and the envelope's first/last ids follow the renamed entries.""" diff --git a/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py index ddac561f337..6a7ec13ec4c 100644 --- a/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py +++ b/tests/unit/llms/anthropic/test_anthropic_count_tokens_transformation.py @@ -1,4 +1,9 @@ +import httpx +import pytest +import respx +import litellm +from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler from litellm.llms.anthropic.count_tokens.transformation import ( AnthropicCountTokensConfig, ) @@ -88,3 +93,72 @@ def test_transform_no_system_no_tools(): assert "system" not in result assert "tools" not in result + + +@pytest.mark.parametrize( + ("api_base", "expected"), + [ + (None, "https://api.anthropic.com/v1/messages/count_tokens"), + ("", "https://api.anthropic.com/v1/messages/count_tokens"), + ("https://gateway.example", "https://gateway.example/v1/messages/count_tokens"), + ("https://gateway.example/", "https://gateway.example/v1/messages/count_tokens"), + ("https://gateway.example/v1", "https://gateway.example/v1/messages/count_tokens"), + ("https://gateway.example/anthropic/v1/messages", "https://gateway.example/anthropic/v1/messages/count_tokens"), + ], +) +def test_endpoint_appends_count_tokens_path_to_deployment_api_base(api_base, expected, monkeypatch): + monkeypatch.delenv("ANTHROPIC_API_BASE", raising=False) + monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False) + assert AnthropicCountTokensConfig().get_anthropic_count_tokens_endpoint(api_base) == expected + + +@pytest.mark.parametrize("env_name", ["ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"]) +@pytest.mark.parametrize("api_base", [None, ""]) +def test_endpoint_without_deployment_api_base_follows_env_base(env_name, api_base, monkeypatch): + """Chat and the federated exchange resolve an unset deployment base through the environment, + so an env-only gateway must receive the count too, never Anthropic's public host.""" + monkeypatch.delenv("ANTHROPIC_API_BASE", raising=False) + monkeypatch.delenv("ANTHROPIC_BASE_URL", raising=False) + monkeypatch.setenv(env_name, "https://env-gateway.example/v1/messages/") + assert ( + AnthropicCountTokensConfig().get_anthropic_count_tokens_endpoint(api_base) + == "https://env-gateway.example/v1/messages/count_tokens" + ) + + +def test_endpoint_prefers_deployment_api_base_over_env_base(monkeypatch): + monkeypatch.setenv("ANTHROPIC_API_BASE", "https://env-gateway.example") + assert ( + AnthropicCountTokensConfig().get_anthropic_count_tokens_endpoint("https://gateway.example/v1") + == "https://gateway.example/v1/messages/count_tokens" + ) + + +@pytest.fixture +def httpx_transport_clients(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + client_cache = getattr(litellm, "in_memory_llm_clients_cache", None) + if client_cache is not None: + client_cache.flush_cache() + yield + if client_cache is not None: + client_cache.flush_cache() + + +@pytest.mark.asyncio +async def test_handler_posts_to_count_tokens_path_under_deployment_api_base(httpx_transport_clients): + """A deployment api_base names the chat host, so a handler that posts to it verbatim lands on + the host root, gets a 404, and the official count silently degrades to the local tokenizer.""" + with respx.mock: + route = respx.post("https://gateway.example/v1/messages/count_tokens").mock( + return_value=httpx.Response(200, json={"input_tokens": 7}) + ) + result = await AnthropicCountTokensHandler().handle_count_tokens_request( + model="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + auth_header={"x-api-key": "sk-ant-api03-test-key"}, + api_base="https://gateway.example", + ) + + assert route.called + assert result == {"input_tokens": 7} diff --git a/tests/unit/llms/anthropic/test_anthropic_files_and_batches.py b/tests/unit/llms/anthropic/test_anthropic_files_and_batches.py index 2728ba03ae4..fe7ea23042f 100644 --- a/tests/unit/llms/anthropic/test_anthropic_files_and_batches.py +++ b/tests/unit/llms/anthropic/test_anthropic_files_and_batches.py @@ -70,9 +70,7 @@ class TestAnthropicFilesHandler: @pytest.fixture def mock_anthropic_batch_results_canceled(self): """Mock Anthropic batch results with canceled status""" - return json.dumps( - {"custom_id": "test-request-3", "result": {"type": "canceled"}} - ).encode("utf-8") + return json.dumps({"custom_id": "test-request-3", "result": {"type": "canceled"}}).encode("utf-8") @pytest.fixture def mock_anthropic_batch_results_mixed(self): @@ -114,9 +112,7 @@ class TestAnthropicFilesHandler: return "\n".join(lines).encode("utf-8") @pytest.mark.asyncio - async def test_afile_content_success( - self, handler, mock_anthropic_batch_results_succeeded - ): + async def test_afile_content_success(self, handler, mock_anthropic_batch_results_succeeded): """Test successful file content retrieval and transformation""" file_content_request: FileContentRequest = { "file_id": "batch_123", @@ -135,16 +131,14 @@ class TestAnthropicFilesHandler: ), ) - with patch( + with patch( # test-quality-ok: the proxy wiring under test is what this patches "litellm.llms.anthropic.files.handler.get_async_httpx_client" - ) as mock_get_client: + ) as mock_get_client: # test-quality-ok: the proxy wiring under test is what this patches mock_client = AsyncMock() mock_client.get = AsyncMock(return_value=mock_response) mock_get_client.return_value = mock_client - with patch.object( - handler.anthropic_model_info, "get_api_key", return_value="test-api-key" - ): + with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"): with patch.object( handler.anthropic_model_info, "get_api_base", @@ -161,9 +155,7 @@ class TestAnthropicFilesHandler: # Verify transformation to OpenAI format content = result.response.content.decode("utf-8") - lines = [ - line for line in content.strip().split("\n") if line.strip() - ] + lines = [line for line in content.strip().split("\n") if line.strip()] assert len(lines) == 1 transformed_result = json.loads(lines[0]) @@ -172,18 +164,13 @@ class TestAnthropicFilesHandler: assert "body" in transformed_result["response"] # Verify body has required OpenAI format fields assert "id" in transformed_result["response"]["body"] - assert ( - transformed_result["response"]["body"]["object"] - == "chat.completion" - ) + assert transformed_result["response"]["body"]["object"] == "chat.completion" assert "choices" in transformed_result["response"]["body"] # Verify request_id matches the original message id assert transformed_result["response"]["request_id"] == "msg_123" @pytest.mark.asyncio - async def test_afile_content_with_prefix( - self, handler, mock_anthropic_batch_results_succeeded - ): + async def test_afile_content_with_prefix(self, handler, mock_anthropic_batch_results_succeeded): """Test file content retrieval with anthropic_batch_results: prefix""" file_content_request: FileContentRequest = { "file_id": "anthropic_batch_results:batch_123", @@ -203,14 +190,12 @@ class TestAnthropicFilesHandler: with patch( "litellm.llms.anthropic.files.handler.get_async_httpx_client" - ) as mock_get_client: + ) as mock_get_client: # test-quality-ok: the proxy wiring under test is what this patches mock_client = AsyncMock() mock_client.get = AsyncMock(return_value=mock_response) mock_get_client.return_value = mock_client - with patch.object( - handler.anthropic_model_info, "get_api_key", return_value="test-api-key" - ): + with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"): with patch.object( handler.anthropic_model_info, "get_api_base", @@ -228,9 +213,7 @@ class TestAnthropicFilesHandler: assert "batch_123" in call_url @pytest.mark.asyncio - async def test_afile_content_errored_result( - self, handler, mock_anthropic_batch_results_errored - ): + async def test_afile_content_errored_result(self, handler, mock_anthropic_batch_results_errored): """Test transformation of errored batch results""" file_content_request: FileContentRequest = { "file_id": "batch_123", @@ -250,14 +233,12 @@ class TestAnthropicFilesHandler: with patch( "litellm.llms.anthropic.files.handler.get_async_httpx_client" - ) as mock_get_client: + ) as mock_get_client: # test-quality-ok: the proxy wiring under test is what this patches mock_client = AsyncMock() mock_client.get = AsyncMock(return_value=mock_response) mock_get_client.return_value = mock_client - with patch.object( - handler.anthropic_model_info, "get_api_key", return_value="test-api-key" - ): + with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"): with patch.object( handler.anthropic_model_info, "get_api_base", @@ -269,29 +250,17 @@ class TestAnthropicFilesHandler: ) content = result.response.content.decode("utf-8") - lines = [ - line for line in content.strip().split("\n") if line.strip() - ] + lines = [line for line in content.strip().split("\n") if line.strip()] assert len(lines) == 1 transformed_result = json.loads(lines[0]) assert transformed_result["custom_id"] == "test-request-2" - assert ( - transformed_result["response"]["status_code"] == 400 - ) # invalid_request_error maps to 400 - assert ( - transformed_result["response"]["body"]["error"]["type"] - == "invalid_request_error" - ) - assert ( - transformed_result["response"]["body"]["error"]["message"] - == "Invalid request" - ) + assert transformed_result["response"]["status_code"] == 400 # invalid_request_error maps to 400 + assert transformed_result["response"]["body"]["error"]["type"] == "invalid_request_error" + assert transformed_result["response"]["body"]["error"]["message"] == "Invalid request" @pytest.mark.asyncio - async def test_afile_content_canceled_result( - self, handler, mock_anthropic_batch_results_canceled - ): + async def test_afile_content_canceled_result(self, handler, mock_anthropic_batch_results_canceled): """Test transformation of canceled batch results""" file_content_request: FileContentRequest = { "file_id": "batch_123", @@ -311,14 +280,12 @@ class TestAnthropicFilesHandler: with patch( "litellm.llms.anthropic.files.handler.get_async_httpx_client" - ) as mock_get_client: + ) as mock_get_client: # test-quality-ok: the proxy wiring under test is what this patches mock_client = AsyncMock() mock_client.get = AsyncMock(return_value=mock_response) mock_get_client.return_value = mock_client - with patch.object( - handler.anthropic_model_info, "get_api_key", return_value="test-api-key" - ): + with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"): with patch.object( handler.anthropic_model_info, "get_api_base", @@ -330,23 +297,16 @@ class TestAnthropicFilesHandler: ) content = result.response.content.decode("utf-8") - lines = [ - line for line in content.strip().split("\n") if line.strip() - ] + lines = [line for line in content.strip().split("\n") if line.strip()] assert len(lines) == 1 transformed_result = json.loads(lines[0]) assert transformed_result["custom_id"] == "test-request-3" assert transformed_result["response"]["status_code"] == 400 - assert ( - "Batch request was canceled" - in transformed_result["response"]["body"]["error"]["message"] - ) + assert "Batch request was canceled" in transformed_result["response"]["body"]["error"]["message"] @pytest.mark.asyncio - async def test_afile_content_mixed_results( - self, handler, mock_anthropic_batch_results_mixed - ): + async def test_afile_content_mixed_results(self, handler, mock_anthropic_batch_results_mixed): """Test transformation of mixed batch results (succeeded, errored, expired)""" file_content_request: FileContentRequest = { "file_id": "batch_123", @@ -366,14 +326,12 @@ class TestAnthropicFilesHandler: with patch( "litellm.llms.anthropic.files.handler.get_async_httpx_client" - ) as mock_get_client: + ) as mock_get_client: # test-quality-ok: the proxy wiring under test is what this patches mock_client = AsyncMock() mock_client.get = AsyncMock(return_value=mock_response) mock_get_client.return_value = mock_client - with patch.object( - handler.anthropic_model_info, "get_api_key", return_value="test-api-key" - ): + with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"): with patch.object( handler.anthropic_model_info, "get_api_base", @@ -385,9 +343,7 @@ class TestAnthropicFilesHandler: ) content = result.response.content.decode("utf-8") - lines = [ - line for line in content.strip().split("\n") if line.strip() - ] + lines = [line for line in content.strip().split("\n") if line.strip()] assert len(lines) == 3 # Check first result (succeeded) @@ -396,9 +352,7 @@ class TestAnthropicFilesHandler: # Check second result (errored) result2 = json.loads(lines[1]) - assert ( - result2["response"]["status_code"] == 429 - ) # rate_limit_error maps to 429 + assert result2["response"]["status_code"] == 429 # rate_limit_error maps to 429 # Check third result (expired) result3 = json.loads(lines[2]) @@ -415,12 +369,12 @@ class TestAnthropicFilesHandler: } with patch.object( - handler.anthropic_model_info, "get_auth_header", return_value=None + handler.anthropic_model_info, + "aget_auth_header", + new=AsyncMock(return_value=None), ): with pytest.raises(ValueError, match="Missing Anthropic API Key"): - await handler.afile_content( - file_content_request=file_content_request, api_key=None - ) + await handler.afile_content(file_content_request=file_content_request, api_key=None) @pytest.mark.asyncio async def test_afile_content_missing_file_id(self, handler): @@ -432,9 +386,7 @@ class TestAnthropicFilesHandler: } with pytest.raises(ValueError, match="file_id is required"): - await handler.afile_content( - file_content_request=file_content_request, api_key="test-api-key" - ) + await handler.afile_content(file_content_request=file_content_request, api_key="test-api-key") @pytest.mark.asyncio async def test_afile_content_http_error(self, handler): @@ -454,21 +406,17 @@ class TestAnthropicFilesHandler: ), ) mock_response.raise_for_status = MagicMock( - side_effect=httpx.HTTPStatusError( - "Not Found", request=mock_response.request, response=mock_response - ) + side_effect=httpx.HTTPStatusError("Not Found", request=mock_response.request, response=mock_response) ) with patch( "litellm.llms.anthropic.files.handler.get_async_httpx_client" - ) as mock_get_client: + ) as mock_get_client: # test-quality-ok: the proxy wiring under test is what this patches mock_client = AsyncMock() mock_client.get = AsyncMock(return_value=mock_response) mock_get_client.return_value = mock_client - with patch.object( - handler.anthropic_model_info, "get_api_key", return_value="test-api-key" - ): + with patch.object(handler.anthropic_model_info, "get_api_key", return_value="test-api-key"): with patch.object( handler.anthropic_model_info, "get_api_base", @@ -480,6 +428,160 @@ class TestAnthropicFilesHandler: api_key="test-api-key", ) + @pytest.mark.asyncio + async def test_afile_content_resolves_wif_via_async_facade( + self, handler, mock_anthropic_batch_results_succeeded, monkeypatch + ): + """Regression: afile_content ran the blocking WIF mint on the event loop + through the sync get_auth_header; it must go through the async facade.""" + import threading + + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import aget_anthropic_wif_token, get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + for name in ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + ): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_files") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-files") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "files-inline-jwt") + + minted = "sk-ant-oat01-files-minted" + thread_ids = [] + + class ThreadRecordingPoster: + def post(self, url, *, content, headers, timeout): + thread_ids.append(threading.get_ident()) + return httpx.Response( + 200, + json={"access_token": minted, "token_type": "Bearer", "expires_in": 3600}, + ) + + engine = JwtBearerTokenExchangeEngine(poster=ThreadRecordingPoster()) + sync_calls = [] + + def sync_shim(litellm_params, api_base, model): + sync_calls.append(model) + return get_anthropic_wif_token(litellm_params, api_base, model, engine) + + async def async_shim(litellm_params, api_base, model): + return await aget_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "get_anthropic_wif_token", sync_shim) + monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim) + + mock_response = httpx.Response( + status_code=200, + content=mock_anthropic_batch_results_succeeded, + headers={"content-type": "application/json"}, + request=httpx.Request( + method="GET", + url="https://api.anthropic.com/v1/messages/batches/batch_123/results", + ), + ) + + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.llms.anthropic.files.handler.get_async_httpx_client" + ) as mock_get_client: # test-quality-ok: the proxy wiring under test is what this patches + mock_client = AsyncMock() + mock_client.get = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + await handler.afile_content( + file_content_request={ + "file_id": "batch_123", + "extra_headers": None, + "extra_body": None, + }, + api_key=None, + ) + + sent_headers = mock_client.get.call_args.kwargs["headers"] + + assert sent_headers["authorization"] == f"Bearer {minted}" + assert "oauth-2025-04-20" in sent_headers["anthropic-beta"] + assert sync_calls == [] + assert thread_ids and thread_ids[0] != threading.get_ident() + + @pytest.mark.asyncio + async def test_afile_content_mints_from_the_deployment_litellm_params( + self, handler, mock_anthropic_batch_results_succeeded, monkeypatch + ): + """Regression: a deployment that authenticates through a named credential carries its + federation settings in litellm_params, and afile_content dropped them, so only + process-wide env vars could ever mint on a batch-result download.""" + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import aget_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + for name in ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + "ANTHROPIC_FEDERATION_RULE_ID", + "ANTHROPIC_ORGANIZATION_ID", + "ANTHROPIC_IDENTITY_TOKEN", + ): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv("CREDENTIAL_IDENTITY_JWT", "credential-inline-jwt") + + minted = "sk-ant-oat01-credential-minted" + + class Poster: + def post(self, url, *, content, headers, timeout): + return httpx.Response( + 200, + json={"access_token": minted, "token_type": "Bearer", "expires_in": 3600}, + ) + + engine = JwtBearerTokenExchangeEngine(poster=Poster()) + + async def async_shim(litellm_params, api_base, model): + return await aget_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim) + + mock_response = httpx.Response( + status_code=200, + content=mock_anthropic_batch_results_succeeded, + headers={"content-type": "application/json"}, + request=httpx.Request( + method="GET", + url="https://api.anthropic.com/v1/messages/batches/batch_123/results", + ), + ) + + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.llms.anthropic.files.handler.get_async_httpx_client" + ) as mock_get_client: # test-quality-ok: the proxy wiring under test is what this patches + mock_client = AsyncMock() + mock_client.get = AsyncMock(return_value=mock_response) + mock_get_client.return_value = mock_client + + await handler.afile_content( + file_content_request={ + "file_id": "batch_123", + "extra_headers": None, + "extra_body": None, + }, + api_key=None, + litellm_params={ + "anthropic_federation_rule_id": "fdrl_credential", + "anthropic_organization_id": "org-credential", + "anthropic_identity_token": "oidc/env/CREDENTIAL_IDENTITY_JWT", + }, + ) + + sent_headers = mock_client.get.call_args.kwargs["headers"] + + assert sent_headers["authorization"] == f"Bearer {minted}" + class TestAnthropicBatchesConfig: """Test Anthropic Batches Config for batch retrieval transformation""" @@ -562,15 +664,11 @@ class TestAnthropicBatchesConfig: ) assert url == "https://api.anthropic.com/v1/messages/batches/batch_123" - def test_transform_retrieve_batch_response_in_progress( - self, config, mock_anthropic_batch_response_in_progress - ): + def test_transform_retrieve_batch_response_in_progress(self, config, mock_anthropic_batch_response_in_progress): """Test transformation of in_progress batch response""" mock_response = httpx.Response( status_code=200, - content=json.dumps(mock_anthropic_batch_response_in_progress).encode( - "utf-8" - ), + content=json.dumps(mock_anthropic_batch_response_in_progress).encode("utf-8"), request=httpx.Request( method="GET", url="https://api.anthropic.com/v1/messages/batches/batch_123", @@ -596,9 +694,7 @@ class TestAnthropicBatchesConfig: assert batch.in_progress_at is not None assert batch.completed_at is None - def test_transform_retrieve_batch_response_completed( - self, config, mock_anthropic_batch_response_completed - ): + def test_transform_retrieve_batch_response_completed(self, config, mock_anthropic_batch_response_completed): """Test transformation of completed batch response""" mock_response = httpx.Response( status_code=200, @@ -624,9 +720,7 @@ class TestAnthropicBatchesConfig: assert batch.request_counts.completed == 10 assert batch.request_counts.failed == 0 - def test_transform_retrieve_batch_response_canceling( - self, config, mock_anthropic_batch_response_canceling - ): + def test_transform_retrieve_batch_response_canceling(self, config, mock_anthropic_batch_response_canceling): """Test transformation of canceling batch response""" mock_response = httpx.Response( status_code=200, @@ -663,9 +757,7 @@ class TestAnthropicBatchesConfig: ) logging_obj = MagicMock() - with pytest.raises( - ValueError, match="Failed to parse Anthropic batch response" - ): + with pytest.raises(ValueError, match="Failed to parse Anthropic batch response"): config.transform_retrieve_batch_response( model="claude-3-5-sonnet-20241022", raw_response=mock_response, diff --git a/tests/unit/llms/anthropic/test_anthropic_prompt_cache_prediction.py b/tests/unit/llms/anthropic/test_anthropic_prompt_cache_prediction.py index 12b81d378c8..9bb26b66aa7 100644 --- a/tests/unit/llms/anthropic/test_anthropic_prompt_cache_prediction.py +++ b/tests/unit/llms/anthropic/test_anthropic_prompt_cache_prediction.py @@ -115,7 +115,16 @@ def test_prediction_header_eligibility(headers: Mapping[str, str], supported: bo @pytest.mark.asyncio -async def test_provider_count_uses_same_version_and_preserves_native_input(monkeypatch: pytest.MonkeyPatch) -> None: +@pytest.mark.parametrize( + "api_base, count_url", + [ + (None, "https://api.anthropic.com/v1/messages/count_tokens"), + ("https://gateway.example/v1/messages", "https://gateway.example/v1/messages/count_tokens"), + ], +) +async def test_provider_count_uses_same_version_and_preserves_native_input( + monkeypatch: pytest.MonkeyPatch, api_base: str | None, count_url: str +) -> None: body: Final = _body() requests: Final[list[httpx.Request]] = [] @@ -128,12 +137,12 @@ async def test_provider_count_uses_same_version_and_preserves_native_input(monke client.client = httpx.AsyncClient(transport=httpx.MockTransport(provider)) monkeypatch.setattr(count_handler, "get_async_httpx_client", lambda **kwargs: client) try: - assert await count_prompt_tokens(_MODEL, _KEY, body) == 311 + assert await count_prompt_tokens(_MODEL, _KEY, body, api_base=api_base) == 311 finally: await client.client.aclose() assert len(requests) == 1 assert requests[0].headers["anthropic-version"] == DEFAULT_ANTHROPIC_API_VERSION - assert requests[0].url == "https://api.anthropic.com/v1/messages/count_tokens" + assert requests[0].url == count_url assert json.loads(requests[0].content) == body diff --git a/tests/unit/llms/anthropic/test_anthropic_wif.py b/tests/unit/llms/anthropic/test_anthropic_wif.py new file mode 100644 index 00000000000..a054b4d130c --- /dev/null +++ b/tests/unit/llms/anthropic/test_anthropic_wif.py @@ -0,0 +1,1293 @@ +import concurrent.futures +import json +from collections.abc import Callable, Mapping +from pathlib import Path +from typing import Final + +import httpx +import jwt +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec + +import litellm +from litellm.llms.anthropic.wif import ( + AnthropicWifParams, + _raise_anthropic_wif_error, + build_anthropic_wif_spec, + get_anthropic_wif_token, + resolve_anthropic_wif_params, +) +from litellm.llms.base_llm.auth.identity_source import ( + InternalIssuerSource, + KeycloakSource, + identity_source_ref, +) +from litellm.llms.base_llm.auth.jwt_signing import build_jwks, rfc7638_thumbprint +from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine +from litellm.llms.base_llm.auth.types import ( + AssertionSourceError, + ExchangeError, + InsecureTokenUrl, + MalformedTokenResponse, + TokenEndpointError, + TokenTransportError, +) +from litellm.types.router import GenericLiteLLMParams + +WIF_ENV_VARS: Final = ( + "ANTHROPIC_FEDERATION_RULE_ID", + "ANTHROPIC_ORGANIZATION_ID", + "ANTHROPIC_SERVICE_ACCOUNT_ID", + "ANTHROPIC_FEDERATION_WORKSPACE_ID", + "ANTHROPIC_IDENTITY_TOKEN_FILE", + "ANTHROPIC_IDENTITY_TOKEN", + "ANTHROPIC_IDENTITY_SOURCE", + "ANTHROPIC_SCOPE", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", +) + +GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer" + + +@pytest.fixture(autouse=True) +def _clean_wif_env(monkeypatch: pytest.MonkeyPatch) -> None: + for name in WIF_ENV_VARS: + monkeypatch.delenv(name, raising=False) + + +class FakeClock: + def __init__(self, start: float = 1_000.0) -> None: + self.now = start + + def __call__(self) -> float: + return self.now + + def advance(self, seconds: float) -> None: + self.now += seconds + + +class RecordedRequest: + def __init__(self, url: str, content: bytes, headers: Mapping[str, str], timeout: float) -> None: + self.url = url + self.content = content + self.headers = dict(headers) + self.timeout = timeout + + def json_body(self) -> dict: + return json.loads(self.content) + + +class ScriptedPoster: + def __init__(self, responses: list[httpx.Response]) -> None: + self.requests: list[RecordedRequest] = [] + self._responses = list(responses) + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.requests.append(RecordedRequest(url, content, headers, timeout)) + if len(self._responses) > 1: + return self._responses.pop(0) + return self._responses[0] + + +class ManualExecutor(concurrent.futures.Executor): + def __init__(self) -> None: + self.pending: list[Callable[[], None]] = [] + + def submit(self, fn, /, *args, **kwargs): + future: concurrent.futures.Future = concurrent.futures.Future() + self.pending.append(lambda: fn(*args, **kwargs)) + return future + + +def token_response(token: str = "sk-ant-oat01-minted", expires_in: int | None = 3600) -> httpx.Response: + body: Final[dict[str, str | int]] = { + "access_token": token, + "token_type": "Bearer", + **({} if expires_in is None else {"expires_in": expires_in}), + } + return httpx.Response(200, json=body) + + +def make_engine(poster: ScriptedPoster, clock: FakeClock | None = None) -> JwtBearerTokenExchangeEngine: + return JwtBearerTokenExchangeEngine( + poster=poster, + clock=clock if clock is not None else FakeClock(), + refresh_executor=ManualExecutor(), + ) + + +def write_token_file(directory: Path, content: str, name: str = "identity-token") -> Path: + directory.mkdir(parents=True, exist_ok=True) + token_file = directory / name + token_file.write_text(content, encoding="utf-8") + return token_file + + +class TestWireProtocolExact: + def test_minimal_body_and_headers(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + monkeypatch.setenv("ANTHROPIC_SCOPE", "user:inference") + token_file = write_token_file(tmp_path, "jwt-assertion-value\n") + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + + token = get_anthropic_wif_token( + { + "anthropic_federation_rule_id": "fdrl_abc123", + "anthropic_organization_id": "org-uuid-1", + "anthropic_identity_token_file": str(token_file), + }, + "https://api.anthropic.com", + "claude-sonnet-4-5", + engine, + ) + + assert token == "sk-ant-oat01-minted" + assert len(poster.requests) == 1 + request = poster.requests[0] + assert request.url == "https://api.anthropic.com/v1/oauth/token" + assert "anthropic-beta" not in request.headers + assert request.headers["content-type"] == "application/json" + assert request.json_body() == { + "grant_type": GRANT_TYPE, + "federation_rule_id": "fdrl_abc123", + "organization_id": "org-uuid-1", + "assertion": "jwt-assertion-value", + } + + def test_optional_fields_present_when_set(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + + get_anthropic_wif_token( + { + "anthropic_federation_rule_id": "fdrl_abc123", + "anthropic_organization_id": "org-uuid-1", + "anthropic_service_account_id": "svcacct_1", + "anthropic_federation_workspace_id": "wrkspc_1", + "anthropic_identity_token_file": str(token_file), + }, + "https://api.anthropic.com", + "claude-sonnet-4-5", + engine, + ) + + request = poster.requests[0] + assert "anthropic-beta" not in request.headers + assert request.headers["content-type"] == "application/json" + assert request.json_body() == { + "grant_type": GRANT_TYPE, + "federation_rule_id": "fdrl_abc123", + "organization_id": "org-uuid-1", + "service_account_id": "svcacct_1", + "workspace_id": "wrkspc_1", + "assertion": "jwt-assertion-value", + } + + def test_spec_cache_key_identity(self): + params = AnthropicWifParams( + federation_rule_id="fdrl_1", + organization_id="org-1", + assertion_ref="oidc/env/ANTHROPIC_IDENTITY_TOKEN", + ) + spec = build_anthropic_wif_spec(params, "https://api.anthropic.com") + assert spec.cache_key_identity == ("fdrl_1", "org-1", "", "") + assert spec.body_encoding == "json" + assert spec.assertion_field == "assertion" + + def test_full_params_spec_has_no_request_headers(self): + """The token exchange sends no anthropic-beta header at all (verified against the + live endpoint); this must hold even for a fully populated params set, so a future + edit cannot reintroduce the header gated on service_account_id or workspace_id.""" + params = AnthropicWifParams( + federation_rule_id="fdrl_1", + organization_id="org-1", + service_account_id="svcacct_1", + workspace_id="wrkspc_1", + assertion_ref="oidc/env/ANTHROPIC_IDENTITY_TOKEN", + ) + spec = build_anthropic_wif_spec(params, "https://api.anthropic.com") + assert dict(spec.request_headers) == {} + + +class TestExchangeHostTrust: + """A federated exchange sends the workload's identity token to api_base and presents the minted + org-scoped token to it, so api_base is a trust decision. Anyone able to write api_base, on the + deployment or on a credential it references, could otherwise redirect both, which is why this is + enforced where the exchange is built rather than at each write path.""" + + def _mint(self, api_base: str | None, monkeypatch: pytest.MonkeyPatch) -> str: + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") + poster = ScriptedPoster([token_response()]) + get_anthropic_wif_token( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, + api_base, + "claude-sonnet-4-5", + make_engine(poster), + ) + return poster.requests[0].url + + def test_anthropic_is_trusted_without_configuration(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", raising=False) + assert self._mint("https://api.anthropic.com", monkeypatch) == "https://api.anthropic.com/v1/oauth/token" + + def test_an_unlisted_host_never_receives_the_identity_token(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", raising=False) + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") + poster = ScriptedPoster([token_response()]) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + get_anthropic_wif_token( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, + "https://attacker.example", + "claude-sonnet-4-5", + make_engine(poster), + ) + + assert poster.requests == [], "the exchange must be refused before anything is sent" + assert "attacker.example" in str(exc_info.value) + assert "LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS" in str(exc_info.value), ( + "an operator running a private gateway has to be told how to allow it" + ) + assert not exc_info.value.message.endswith(".") + + def test_a_lookalike_host_does_not_pass_on_a_substring(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", raising=False) + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") + poster = ScriptedPoster([token_response()]) + + with pytest.raises(litellm.AuthenticationError): + get_anthropic_wif_token( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, + "https://api.anthropic.com.evil.test", + "claude-sonnet-4-5", + make_engine(poster), + ) + + assert poster.requests == [] + + def test_an_operator_can_allow_a_private_gateway(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal") + assert self._mint("https://gateway.internal", monkeypatch) == "https://gateway.internal/v1/oauth/token" + + def test_a_gateway_listed_with_its_port_is_trusted(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal:8443") + assert self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" + + def test_allowlist_matching_ignores_hostname_case(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "Gateway.Internal:8443") + assert self._mint("https://gateway.internal:8443", monkeypatch) == "https://gateway.internal:8443/v1/oauth/token" + assert self._mint("https://GATEWAY.internal:8443", monkeypatch) == "https://GATEWAY.internal:8443/v1/oauth/token" + + def test_a_gateway_listed_with_a_port_is_not_trusted_on_another_port(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal:8443") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") + poster = ScriptedPoster([token_response()]) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + get_anthropic_wif_token( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, + "https://gateway.internal:9443", + "claude-sonnet-4-5", + make_engine(poster), + ) + + assert poster.requests == [], "another process on the same host is not the allowed gateway" + assert "gateway.internal:9443" in str(exc_info.value) + + def test_a_gateway_listed_without_a_port_is_trusted_on_every_port(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal") + assert self._mint("https://gateway.internal:9443", monkeypatch) == "https://gateway.internal:9443/v1/oauth/token" + + def test_an_entry_spelling_the_scheme_default_port_matches_a_base_that_omits_it( + self, monkeypatch: pytest.MonkeyPatch + ): + monkeypatch.setenv("LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", "gateway.internal:443") + assert self._mint("https://gateway.internal", monkeypatch) == "https://gateway.internal/v1/oauth/token" + + + +class TestBaseUrlDerivation: + def _mint(self, api_base: str | None, monkeypatch: pytest.MonkeyPatch) -> str: + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") + # These cases are about how a base is normalised into a token URL, not about which hosts an + # operator trusts, so the private hosts they use are allowlisted explicitly. The trust + # boundary itself is covered by TestExchangeHostTrust. + monkeypatch.setenv( + "LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS", + "gw.example.com,env.example.com,base.example.com,model.example.com", + ) + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + get_anthropic_wif_token( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, + api_base, + "claude-sonnet-4-5", + engine, + ) + return poster.requests[0].url + + def test_explicit_api_base_wins(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_API_BASE", "https://env.example.com") + assert self._mint("https://gw.example.com/", monkeypatch) == "https://gw.example.com/v1/oauth/token" + + def test_env_api_base(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_API_BASE", "https://env.example.com") + assert self._mint(None, monkeypatch) == "https://env.example.com/v1/oauth/token" + + def test_empty_api_base_falls_back_like_unset(self, monkeypatch: pytest.MonkeyPatch): + """Chat treats an empty deployment api_base as unset; the exchange must not refuse host ''.""" + monkeypatch.setenv("ANTHROPIC_API_BASE", "https://env.example.com") + assert self._mint("", monkeypatch) == "https://env.example.com/v1/oauth/token" + + def test_env_base_url(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_BASE_URL", "https://base.example.com") + assert self._mint(None, monkeypatch) == "https://base.example.com/v1/oauth/token" + + def test_default_base(self, monkeypatch: pytest.MonkeyPatch): + assert self._mint(None, monkeypatch) == "https://api.anthropic.com/v1/oauth/token" + + @pytest.mark.parametrize( + "api_base", + [ + "https://gw.example.com/v1/messages", + "https://gw.example.com/v1/messages/", + "https://gw.example.com/v1/messages//v1/messages", + ], + ) + def test_chat_appended_bases_normalize_to_clean_token_url(self, api_base: str, monkeypatch: pytest.MonkeyPatch): + """main.py appends /v1/messages before dispatch (twice for trailing-slash + bases); the exchange must still target the deployment base.""" + assert self._mint(api_base, monkeypatch) == "https://gw.example.com/v1/oauth/token" + + def test_trailing_slash_env_base_normalizes(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_BASE_URL", "https://base.example.com/") + assert self._mint(None, monkeypatch) == "https://base.example.com/v1/oauth/token" + + +class TestSecretManagerEnvResolution: + """WIF env vars resolve through get_secret_str so configured secret managers + work, exactly like every sibling Anthropic credential.""" + + def test_values_resolve_through_get_secret_str(self, monkeypatch: pytest.MonkeyPatch): + secrets: Final = { + "ANTHROPIC_FEDERATION_RULE_ID": "fdrl_sm", + "ANTHROPIC_ORGANIZATION_ID": "org-sm", + "ANTHROPIC_IDENTITY_TOKEN": "sm-inline-jwt", + } + monkeypatch.setattr( + "litellm.secret_managers.main.get_secret_str", + lambda secret_name, default_value=None: secrets.get(secret_name, default_value), + ) + + params = resolve_anthropic_wif_params(None) + + assert params == AnthropicWifParams( + federation_rule_id="fdrl_sm", + organization_id="org-sm", + assertion_ref="oidc/env/ANTHROPIC_IDENTITY_TOKEN", + ) + + def test_non_str_secret_value_treated_as_unset(self, monkeypatch: pytest.MonkeyPatch): + secrets: Final = { + "ANTHROPIC_FEDERATION_RULE_ID": {"unexpected": "shape"}, + "ANTHROPIC_ORGANIZATION_ID": "org-sm", + "ANTHROPIC_IDENTITY_TOKEN": "sm-inline-jwt", + } + monkeypatch.setattr( + "litellm.secret_managers.main.get_secret_str", + lambda secret_name, default_value=None: secrets.get(secret_name, default_value), + ) + + assert resolve_anthropic_wif_params(None) is None + + +class TestResolutionMatrix: + def test_params_beat_env_per_field(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env") + monkeypatch.setenv("ANTHROPIC_SERVICE_ACCOUNT_ID", "svc-env") + monkeypatch.setenv("ANTHROPIC_FEDERATION_WORKSPACE_ID", "wrkspc_env") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", "/var/run/secrets/env-token") + + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_param", + "anthropic_organization_id": "org-param", + "anthropic_service_account_id": "svc-param", + "anthropic_federation_workspace_id": "wrkspc_param", + "anthropic_identity_token_file": "/var/run/secrets/param-token", + } + ) + + assert params == AnthropicWifParams( + federation_rule_id="fdrl_param", + organization_id="org-param", + service_account_id="svc-param", + workspace_id="wrkspc_param", + assertion_ref="oidc/file//var/run/secrets/param-token", + ) + + def test_env_only_config(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "raw-env-jwt") + + params = resolve_anthropic_wif_params(None) + + assert params is not None + assert params.assertion_ref == "oidc/env/ANTHROPIC_IDENTITY_TOKEN" + assert params.service_account_id is None + assert params.workspace_id is None + + def test_file_param_beats_inline_param(self): + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token_file": "/var/run/secrets/tok", + "anthropic_identity_token": "oidc/env/OTHER", + } + ) + assert params is not None + assert params.assertion_ref == "oidc/file//var/run/secrets/tok" + + def test_inline_param_beats_env_file(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", "/var/run/secrets/env-tok") + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token": "oidc/env/OTHER", + } + ) + assert params is not None + assert params.assertion_ref == "oidc/env/OTHER" + + def test_env_file_beats_env_inline(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", "/var/run/secrets/env-tok") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "raw-env-jwt") + params = resolve_anthropic_wif_params( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"} + ) + assert params is not None + assert params.assertion_ref == "oidc/file//var/run/secrets/env-tok" + + def test_param_token_ref_beats_env_identity_source(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_SOURCE", "internal_issuer") + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token_file": "/var/run/secrets/dep-tok", + } + ) + assert params is not None + assert params.assertion_ref == "oidc/file//var/run/secrets/dep-tok" + assert params.assertion_source is None + + def test_param_inline_token_beats_env_identity_source(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_SOURCE", "internal_issuer") + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token": "oidc/env/OTHER", + } + ) + assert params is not None + assert params.assertion_ref == "oidc/env/OTHER" + assert params.assertion_source is None + + def test_env_identity_source_beats_env_token_refs(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_SOURCE", "internal_issuer") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN_FILE", "/var/run/secrets/env-tok") + with pytest.raises(litellm.AuthenticationError): + resolve_anthropic_wif_params( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"} + ) + + def test_env_identity_source_dispatches_param_issuer_fields(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_SOURCE", "internal_issuer") + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + } + ) + assert params is not None + assert params.assertion_ref.startswith("oidc/internal_issuer/") + assert params.assertion_source is not None + + def test_empty_workspace_env_coerced_to_none(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_FEDERATION_WORKSPACE_ID", "") + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token": "oidc/env/TOK", + } + ) + assert params is not None + assert params.workspace_id is None + spec = build_anthropic_wif_spec(params, "https://api.anthropic.com") + assert "workspace_id" not in spec.static_body + + @pytest.mark.parametrize( + "litellm_params", + [ + {}, + {"anthropic_federation_rule_id": "fdrl_1"}, + {"anthropic_organization_id": "org-1"}, + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, + {"anthropic_organization_id": "org-1", "anthropic_identity_token": "oidc/env/TOK"}, + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_identity_token": "oidc/env/TOK"}, + ], + ) + def test_gate_unmet_returns_none(self, litellm_params: dict): + assert resolve_anthropic_wif_params(litellm_params) is None + + def test_gate_unmet_facade_returns_none_without_engine_call(self): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + assert get_anthropic_wif_token({}, None, "claude-sonnet-4-5", engine) is None + assert poster.requests == [] + + +class TestServiceAccountIdIsOptional: + """Anthropic's reference docs mark service_account_id required, but a live exchange + against a federation rule targeting a single service account mints successfully + without it; resolution must not gate activation on it, and the wire body must omit + the key entirely rather than send it as null.""" + + def test_activates_and_omits_service_account_id_when_unset(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + litellm_params: Final = { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token_file": str(token_file), + } + + params = resolve_anthropic_wif_params(litellm_params) + assert params is not None + assert params.service_account_id is None + + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + token = get_anthropic_wif_token(litellm_params, "https://api.anthropic.com", "claude-sonnet-4-5", engine) + + assert token == "sk-ant-oat01-minted" + assert "service_account_id" not in poster.requests[0].json_body() + + +class TestInlineRefRestrictions: + RAW_JWT: Final = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ3b3JrbG9hZCJ9.c2lnbmF0dXJl" + + @pytest.mark.parametrize("bad_ref", [RAW_JWT, "oidc/env_path/ANTHROPIC_TOKEN_PATH"]) + def test_rejected_inline_refs(self, bad_ref: str): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + get_anthropic_wif_token( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token": bad_ref, + }, + None, + "claude-sonnet-4-5", + engine, + ) + + assert "oidc/env/" in exc_info.value.message + assert "oidc/file/" in exc_info.value.message + assert self.RAW_JWT not in exc_info.value.message + assert poster.requests == [] + + +class TestFileAllowlistAndSymlink: + SECRET_CONTENT: Final = "super-secret-jwt-content" + + def _call(self, token_file: Path, poster: ScriptedPoster) -> str | None: + engine = make_engine(poster) + return get_anthropic_wif_token( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token_file": str(token_file), + }, + "https://api.anthropic.com", + "claude-sonnet-4-5", + engine, + ) + + def test_file_outside_allowlist_rejected(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path / "allowed")) + token_file = write_token_file(tmp_path / "outside", self.SECRET_CONTENT) + poster = ScriptedPoster([token_response()]) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + self._call(token_file, poster) + + assert str(token_file) in exc_info.value.message + assert self.SECRET_CONTENT not in exc_info.value.message + assert poster.requests == [] + + def test_disallowed_path_message_names_allowlist_and_env_var(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + """The disallowed_path error must explain the allowlist and name the env var an + operator would set, not surface as a bare '(disallowed_path)' code dump.""" + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path / "allowed")) + token_file = write_token_file(tmp_path / "outside", self.SECRET_CONTENT) + poster = ScriptedPoster([token_response()]) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + self._call(token_file, poster) + + message = exc_info.value.message + assert "(disallowed_path)" not in message + assert "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS" in message + assert "allowed credential director" in message + + def test_symlink_escape_rejected(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + allowed = tmp_path / "allowed" + allowed.mkdir() + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(allowed)) + outside_file = write_token_file(tmp_path / "outside", self.SECRET_CONTENT) + link = allowed / "identity-token" + link.symlink_to(outside_file) + poster = ScriptedPoster([token_response()]) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + self._call(link, poster) + + assert self.SECRET_CONTENT not in exc_info.value.message + assert poster.requests == [] + + def test_file_inside_allowlist_succeeds(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, self.SECRET_CONTENT) + poster = ScriptedPoster([token_response()]) + + assert self._call(token_file, poster) == "sk-ant-oat01-minted" + assert poster.requests[0].json_body()["assertion"] == self.SECRET_CONTENT + + +class TestErrorMappingExhaustive: + @pytest.mark.parametrize( + "error", + [ + AssertionSourceError(kind="missing", source_ref="oidc/env/TOK"), + AssertionSourceError(kind="disallowed_path", source_ref="oidc/file//etc/passwd"), + InsecureTokenUrl(host="token.example"), + TokenEndpointError(status_code=500, redacted_body="error: server_error"), + TokenTransportError(detail="ConnectError: refused"), + MalformedTokenResponse(detail="token response failed RFC 6749 5.1 schema validation"), + ], + ) + def test_every_variant_maps_to_authentication_error(self, error: ExchangeError): + with pytest.raises(litellm.AuthenticationError) as exc_info: + _raise_anthropic_wif_error( + error, model="claude-sonnet-4-5", workspace_id_set=False, service_account_id_set=False + ) + + assert exc_info.value.llm_provider == "anthropic" + assert exc_info.value.model == "claude-sonnet-4-5" + assert not exc_info.value.message.endswith(".") + + def test_assertion_source_error_detail_is_rendered_when_present(self): + with pytest.raises(litellm.AuthenticationError) as exc_info: + _raise_anthropic_wif_error( + AssertionSourceError(kind="unreadable", source_ref="oidc/keycloak/abc123", detail="invalid_client"), + model="claude-sonnet-4-5", + workspace_id_set=True, + service_account_id_set=True, + ) + + assert "invalid_client" in exc_info.value.message + + def test_assertion_source_error_without_detail_is_unchanged(self): + """Regression floor: the token_file/env path never populates detail, so nothing follows the + source ref and the message ends without a period for the router's suffix.""" + with pytest.raises(litellm.AuthenticationError) as exc_info: + _raise_anthropic_wif_error( + AssertionSourceError(kind="unreadable", source_ref="oidc/env/ANTHROPIC_IDENTITY_TOKEN"), + model="claude-sonnet-4-5", + workspace_id_set=True, + service_account_id_set=True, + ) + + assert exc_info.value.message == ( + "litellm.AuthenticationError: Anthropic workload identity federation failed. Could not obtain " + "the OIDC identity token (unreadable) from oidc/env/ANTHROPIC_IDENTITY_TOKEN" + ) + + def test_endpoint_error_raised_through_facade(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") + poster = ScriptedPoster([httpx.Response(500, json={"error": "server_error"})]) + engine = make_engine(poster) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + get_anthropic_wif_token( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, + None, + "claude-sonnet-4-5", + engine, + ) + + assert exc_info.value.llm_provider == "anthropic" + assert "HTTP 500" in exc_info.value.message + assert "server_error" in exc_info.value.message + + @pytest.mark.parametrize( + "litellm_params,status_code,body", + [ + ( + {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"}, + 401, + {"error": "invalid_grant."}, + ), + ( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_federation_workspace_id": "wrkspc_1", + }, + 500, + {"error": "server_error."}, + ), + ], + ) + def test_token_endpoint_error_message_has_no_doubled_period( + self, litellm_params: dict, status_code: int, body: dict, monkeypatch: pytest.MonkeyPatch + ): + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") + poster = ScriptedPoster([httpx.Response(status_code, json=body)]) + engine = make_engine(poster) + + with pytest.raises(litellm.AuthenticationError) as exc_info: + get_anthropic_wif_token(litellm_params, None, "claude-sonnet-4-5", engine) + + assert ".." not in exc_info.value.message + + +class TestDenialHints: + """Anthropic answers every denied exchange with an opaque 401 and logs the reason + (workspace_id_required, jti_reused, ...) only in the Console, so the error must say where + to look and name whichever optional id is still unset.""" + + BASE_PARAMS: Final = {"anthropic_federation_rule_id": "fdrl_1", "anthropic_organization_id": "org-1"} + + def _raise(self, litellm_params: dict, status_code: int, monkeypatch: pytest.MonkeyPatch) -> str: + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "inline-jwt") + poster = ScriptedPoster([httpx.Response(status_code, json={"error": "invalid_grant"})]) + engine = make_engine(poster) + with pytest.raises(litellm.AuthenticationError) as exc_info: + get_anthropic_wif_token(litellm_params, None, "claude-sonnet-4-5", engine) + return exc_info.value.message + + def test_401_points_at_console_authentication_history(self, monkeypatch: pytest.MonkeyPatch): + message = self._raise(self.BASE_PARAMS, 401, monkeypatch) + assert "authentication history" in message + assert "workspace_id_required" in message + + def test_500_carries_no_denial_hints(self, monkeypatch: pytest.MonkeyPatch): + message = self._raise(self.BASE_PARAMS, 500, monkeypatch) + assert "authentication history" not in message + assert "ANTHROPIC_FEDERATION_WORKSPACE_ID" not in message + assert "ANTHROPIC_SERVICE_ACCOUNT_ID" not in message + + def test_hints_name_both_ids_when_both_unset(self, monkeypatch: pytest.MonkeyPatch): + message = self._raise(self.BASE_PARAMS, 401, monkeypatch) + assert "anthropic_federation_workspace_id" in message + assert "ANTHROPIC_FEDERATION_WORKSPACE_ID" in message + assert "anthropic_service_account_id" in message + assert "ANTHROPIC_SERVICE_ACCOUNT_ID" in message + assert not message.endswith(".") + + def test_the_workspace_hint_says_federation_ignores_the_bedrock_variable(self, monkeypatch: pytest.MonkeyPatch): + """ANTHROPIC_WORKSPACE_ID is the spelling Anthropic's own reference uses, and the Bedrock Claude + platform provider already reads it, so an operator who set it needs the 401 to say it is ignored + here rather than name only a variable they have never heard of.""" + message = self._raise(self.BASE_PARAMS, 401, monkeypatch) + assert "ANTHROPIC_WORKSPACE_ID" in message + assert "Bedrock" in message + + def test_no_workspace_hint_when_workspace_set(self, monkeypatch: pytest.MonkeyPatch): + message = self._raise({**self.BASE_PARAMS, "anthropic_federation_workspace_id": "wrkspc_1"}, 401, monkeypatch) + assert "ANTHROPIC_FEDERATION_WORKSPACE_ID" not in message + assert "ANTHROPIC_SERVICE_ACCOUNT_ID" in message + + def test_no_service_account_hint_when_service_account_set(self, monkeypatch: pytest.MonkeyPatch): + message = self._raise({**self.BASE_PARAMS, "anthropic_service_account_id": "svac_1"}, 401, monkeypatch) + assert "ANTHROPIC_SERVICE_ACCOUNT_ID" not in message + assert "ANTHROPIC_FEDERATION_WORKSPACE_ID" in message + + def test_only_console_pointer_when_both_set(self, monkeypatch: pytest.MonkeyPatch): + message = self._raise( + {**self.BASE_PARAMS, "anthropic_federation_workspace_id": "wrkspc_1", "anthropic_service_account_id": "svac_1"}, + 401, + monkeypatch, + ) + assert "authentication history" in message + assert "ANTHROPIC_FEDERATION_WORKSPACE_ID" not in message + assert "ANTHROPIC_SERVICE_ACCOUNT_ID" not in message + assert ".." not in message + assert not message.endswith(".") + + +class TestFileRereadOnRefresh: + def test_mandatory_refresh_carries_rotated_assertion(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, "first-assertion") + clock = FakeClock(start=1_000.0) + poster = ScriptedPoster( + [token_response("sk-ant-oat01-first", 3600), token_response("sk-ant-oat01-second", 3600)] + ) + engine = make_engine(poster, clock=clock) + litellm_params = { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token_file": str(token_file), + } + + first = get_anthropic_wif_token(litellm_params, "https://api.anthropic.com", "claude-sonnet-4-5", engine) + token_file.write_text("second-assertion", encoding="utf-8") + clock.advance(3600 - 10) + second = get_anthropic_wif_token(litellm_params, "https://api.anthropic.com", "claude-sonnet-4-5", engine) + + assert first == "sk-ant-oat01-first" + assert second == "sk-ant-oat01-second" + assert len(poster.requests) == 2 + assert poster.requests[1].json_body()["assertion"] == "second-assertion" + + +_ISSUER_PRIVATE_VALUE: Final = 55566677788899900011122233344455566677788899900011122233344455 +ISSUER_SIGNING_KEY_REF: Final = "oidc/env/ISSUER_SIGNING_KEY_PEM" +KEYCLOAK_TOKEN_URL: Final = "https://keycloak.internal.example/realms/litellm/protocol/openid-connect/token" + + +def _issuer_signing_key() -> ec.EllipticCurvePrivateKey: + return ec.derive_private_key(_ISSUER_PRIVATE_VALUE, ec.SECP256R1()) + + +def _issuer_signing_key_pem() -> str: + return ( + _issuer_signing_key() + .private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + .decode() + ) + + +def _get_secret_str_returning(pem: str, ref: str) -> Callable[..., str | None]: + def fake_get_secret_str(secret_name: str, default_value: str | None = None) -> str | None: + return pem if secret_name == ref else default_value + + return fake_get_secret_str + + +class TestIdentitySourceDiscriminatorAbsentIsByteIdenticalToLegacy: + """anthropic_identity_source unset must resolve exactly like today: no new dispatch code + runs, and no assertion_source closure is attached, so the engine falls back to its own + reader precisely as it always has.""" + + def test_file_config_carries_no_assertion_source(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + + params = resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_token_file": str(token_file), + } + ) + + assert params == AnthropicWifParams( + federation_rule_id="fdrl_1", + organization_id="org-1", + assertion_ref=f"oidc/file/{token_file}", + ) + assert params.assertion_source is None + + def test_env_config_carries_no_assertion_source(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "raw-env-jwt") + + params = resolve_anthropic_wif_params(None) + + assert params is not None + assert params.assertion_ref == "oidc/env/ANTHROPIC_IDENTITY_TOKEN" + assert params.assertion_source is None + + +class TestInternalIssuerIdentitySourceDispatch: + """A config.yaml-shaped litellm_params block for the internal_issuer identity source.""" + + LITELLM_PARAMS: Final = { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_ttl_seconds": 300, + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + } + + def test_assertion_ref_matches_the_identity_source_hash(self): + params = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + + assert params is not None + expected_config = InternalIssuerSource( + issuer_url="https://issuer.internal.example", + subject="workload-a", + ttl_seconds=300, + signing_key_ref=ISSUER_SIGNING_KEY_REF, + ) + assert params.assertion_ref == identity_source_ref(expected_config) + assert params.assertion_ref.startswith("oidc/internal_issuer/") + + def test_ref_is_stable_and_rolls_on_field_change(self): + first = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + second = resolve_anthropic_wif_params(dict(self.LITELLM_PARAMS)) + changed = resolve_anthropic_wif_params({**self.LITELLM_PARAMS, "anthropic_issuer_subject": "workload-b"}) + + assert first is not None and second is not None and changed is not None + assert first.assertion_ref == second.assertion_ref + assert first.assertion_ref != changed.assertion_ref + + def test_assertion_source_mints_a_verifiable_jwt(self, monkeypatch: pytest.MonkeyPatch): + pem = _issuer_signing_key_pem() + monkeypatch.setattr( + "litellm.secret_managers.main.get_secret_str", + _get_secret_str_returning(pem, ISSUER_SIGNING_KEY_REF), + ) + + params = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + assert params is not None + assert params.assertion_source is not None + + assertion = params.assertion_source() + + assert assertion is not None + public_key = _issuer_signing_key().public_key() + expected_kid = build_jwks(public_key)["keys"][0]["kid"] + assert jwt.get_unverified_header(assertion)["kid"] == expected_kid + assert expected_kid == rfc7638_thumbprint(public_key) + claims = jwt.decode(assertion, public_key, algorithms=["ES256"], options={"verify_aud": False}) + assert claims["sub"] == "workload-a" + assert claims["iss"] == "https://issuer.internal.example" + + def test_full_exchange_sends_the_minted_assertion(self, monkeypatch: pytest.MonkeyPatch): + pem = _issuer_signing_key_pem() + monkeypatch.setattr( + "litellm.secret_managers.main.get_secret_str", + _get_secret_str_returning(pem, ISSUER_SIGNING_KEY_REF), + ) + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + + token = get_anthropic_wif_token(self.LITELLM_PARAMS, "https://api.anthropic.com", "claude-sonnet-4-5", engine) + + assert token == "sk-ant-oat01-minted" + sent_assertion = poster.requests[0].json_body()["assertion"] + jwt.decode( + sent_assertion, _issuer_signing_key().public_key(), algorithms=["ES256"], options={"verify_aud": False} + ) + + +class TestKeycloakIdentitySourceDispatch: + """A config.yaml-shaped litellm_params block for the keycloak identity source. The minted + closure's own network behavior is covered by test_client_credentials.py's DI-poster tests; + this only proves wif.py threads the fields into the right config and hash.""" + + LITELLM_PARAMS: Final = { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": KEYCLOAK_TOKEN_URL, + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_client_secret_ref": "oidc/env/KEYCLOAK_CLIENT_SECRET", + } + + def test_assertion_ref_matches_the_identity_source_hash(self): + params = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + + assert params is not None + expected_config = KeycloakSource( + token_url=KEYCLOAK_TOKEN_URL, + client_id="litellm", + client_secret_ref="oidc/env/KEYCLOAK_CLIENT_SECRET", + ) + assert params.assertion_ref == identity_source_ref(expected_config) + assert params.assertion_ref.startswith("oidc/keycloak/") + + def test_assertion_source_is_a_fresh_closure(self): + params = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + + assert params is not None + assert params.assertion_source is not None + assert callable(params.assertion_source) + + def test_auth_method_change_rolls_the_ref(self): + default_method = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + post_method = resolve_anthropic_wif_params( + {**self.LITELLM_PARAMS, "anthropic_keycloak_auth_method": "client_secret_post"} + ) + + assert default_method is not None and post_method is not None + assert default_method.assertion_ref != post_method.assertion_ref + + def test_client_secret_ref_pointer_name_change_rolls_the_ref_without_resolving_it(self): + """The hash covers the pointer NAME, never a resolved secret (decision 7) -- true even + though nothing in this test ever calls get_secret_str.""" + first = resolve_anthropic_wif_params(self.LITELLM_PARAMS) + second = resolve_anthropic_wif_params( + {**self.LITELLM_PARAMS, "anthropic_keycloak_client_secret_ref": "oidc/env/OTHER_SECRET_NAME"} + ) + + assert first is not None and second is not None + assert first.assertion_ref != second.assertion_ref + + + +@pytest.mark.parametrize( + "sparse_params", + [TestInternalIssuerIdentitySourceDispatch.LITELLM_PARAMS, TestKeycloakIdentitySourceDispatch.LITELLM_PARAMS], + ids=["internal_issuer", "keycloak"], +) +def test_dense_router_params_dump_resolves_like_the_sparse_config(sparse_params: Mapping[str, object]): + dense_params = dict(GenericLiteLLMParams(**sparse_params)) + assert any(value is None for value in dense_params.values()) + + dense = resolve_anthropic_wif_params(dense_params) + sparse = resolve_anthropic_wif_params(sparse_params) + + assert dense is not None and sparse is not None + assert dense.assertion_ref == sparse.assertion_ref + + +class TestIdentitySourceValidationFailsClosed: + """Unknown discriminator, a missing required variant field, and a field belonging to the + other variant are all hard config errors at resolution time -- never a silent fallback to + token_file (decision 5).""" + + def test_unknown_discriminator_raises(self): + with pytest.raises(litellm.AuthenticationError, match="anthropic_identity_source"): + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "bogus", + } + ) + + def test_internal_issuer_missing_required_fields_raises(self): + with pytest.raises(litellm.AuthenticationError): + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + } + ) + + def test_keycloak_missing_required_fields_raises(self): + with pytest.raises(litellm.AuthenticationError): + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_client_id": "litellm", + } + ) + + def test_mixed_variant_fields_raise(self): + with pytest.raises(litellm.AuthenticationError, match="belongs to a different identity source"): + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + "anthropic_keycloak_client_id": "leaked-from-other-variant", + } + ) + + def test_blank_optional_and_foreign_fields_count_as_unset(self): + configured = { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + } + with_blanks = { + **configured, + "anthropic_issuer_audience": "", + "anthropic_issuer_ttl_seconds": "", + "anthropic_keycloak_client_id": "", + } + + expected = resolve_anthropic_wif_params(configured) + actual = resolve_anthropic_wif_params(with_blanks) + + assert expected is not None and actual is not None + assert actual.assertion_ref == expected.assertion_ref + + def test_secret_pasted_into_wrong_field_never_appears_in_the_error(self): + secret_value = "super-secret-client-value-xyz" + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params( + { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_organization_id": "org-1", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + "anthropic_issuer_ttl_seconds": secret_value, + } + ) + + assert secret_value not in exc_info.value.message + + +class TestMissingIdsFailClosedWhenIdentitySourceConfigured: + """An explicit identity source is a request to federate. Without the rule or organization id + the exchange cannot even be attempted, so resolution must say which ids are missing instead + of returning None and letting the request die later as a missing API key.""" + + INTERNAL_ISSUER_FIELDS: Final = { + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.internal.example", + "anthropic_issuer_subject": "workload-a", + "anthropic_issuer_signing_key_ref": ISSUER_SIGNING_KEY_REF, + } + + def test_both_ids_missing_names_both(self): + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params(self.INTERNAL_ISSUER_FIELDS) + + message = exc_info.value.message + assert "'internal_issuer'" in message + assert "anthropic_federation_rule_id and anthropic_organization_id are not set" in message + assert "Settings > Workload identity" in message + assert "ANTHROPIC_FEDERATION_RULE_ID" in message + assert not message.endswith(".") + + def test_only_rule_id_missing_names_only_the_rule(self): + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params({**self.INTERNAL_ISSUER_FIELDS, "anthropic_organization_id": "org-1"}) + + assert "but anthropic_federation_rule_id is not set" in exc_info.value.message + + def test_only_organization_id_missing_names_only_the_org(self): + with pytest.raises(litellm.AuthenticationError) as exc_info: + resolve_anthropic_wif_params({**self.INTERNAL_ISSUER_FIELDS, "anthropic_federation_rule_id": "fdrl_1"}) + + assert "but anthropic_organization_id is not set" in exc_info.value.message + + def test_keycloak_source_fails_closed_too(self): + with pytest.raises(litellm.AuthenticationError, match="'keycloak', but anthropic_federation_rule_id"): + resolve_anthropic_wif_params( + {"anthropic_identity_source": "keycloak", "anthropic_organization_id": "org-1"} + ) + + def test_env_configured_source_fails_closed(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_SOURCE", "internal_issuer") + with pytest.raises(litellm.AuthenticationError, match="anthropic_organization_id is not set"): + resolve_anthropic_wif_params({"anthropic_federation_rule_id": "fdrl_1"}) + + def test_env_ids_satisfy_the_gate(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_env") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-env") + params = resolve_anthropic_wif_params(self.INTERNAL_ISSUER_FIELDS) + assert params is not None + assert params.federation_rule_id == "fdrl_env" + + def test_unknown_source_with_missing_ids_reports_the_unknown_source(self): + with pytest.raises(litellm.AuthenticationError, match="must be one of internal_issuer, keycloak"): + resolve_anthropic_wif_params({"anthropic_identity_source": "bogus"}) + + def test_legacy_token_params_without_ids_still_return_none(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("ANTHROPIC_IDENTITY_SOURCE", "internal_issuer") + assert resolve_anthropic_wif_params({"anthropic_identity_token": "oidc/env/TOK"}) is None + + +class TestConfigYamlShapedIdentitySources: + """One litellm_params dict per identity source, shaped exactly like the + model_list[].litellm_params block a proxy config.yaml carries -- proving an operator can + configure each of Phase 1's supported sources.""" + + def test_legacy_token_file_source(self, tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", str(tmp_path)) + token_file = write_token_file(tmp_path, "jwt-assertion-value") + litellm_params = { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_prod", + "anthropic_organization_id": "org_prod", + "anthropic_identity_token_file": str(token_file), + } + + params = resolve_anthropic_wif_params(litellm_params) + + assert params is not None + assert params.assertion_ref == f"oidc/file/{token_file}" + assert params.assertion_source is None + + def test_internal_issuer_source(self): + litellm_params = { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_prod", + "anthropic_organization_id": "org_prod", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://litellm.internal.example", + "anthropic_issuer_subject": "litellm-proxy", + "anthropic_issuer_ttl_seconds": 300, + "anthropic_issuer_signing_key_ref": "os.environ/ISSUER_SIGNING_KEY_PEM", + } + + params = resolve_anthropic_wif_params(litellm_params) + + assert params is not None + assert params.assertion_ref.startswith("oidc/internal_issuer/") + assert params.assertion_source is not None + + def test_keycloak_source(self): + litellm_params = { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_prod", + "anthropic_organization_id": "org_prod", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": KEYCLOAK_TOKEN_URL, + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_auth_method": "client_secret_post", + "anthropic_keycloak_client_secret_ref": "os.environ/KEYCLOAK_CLIENT_SECRET", + "anthropic_keycloak_scope": "anthropic-wif", + } + + params = resolve_anthropic_wif_params(litellm_params) + + assert params is not None + assert params.assertion_ref.startswith("oidc/keycloak/") + assert params.assertion_source is not None diff --git a/tests/unit/llms/anthropic/test_cost_calculation_dict_safety.py b/tests/unit/llms/anthropic/test_cost_calculation_dict_safety.py index 44b8bb3c9a2..58018a665bc 100644 --- a/tests/unit/llms/anthropic/test_cost_calculation_dict_safety.py +++ b/tests/unit/llms/anthropic/test_cost_calculation_dict_safety.py @@ -5,7 +5,6 @@ being either a ``dict`` or a ``ServerToolUse`` pydantic instance. See https://github.com/BerriAI/litellm/issues/26153. """ - import pytest from litellm.litellm_core_utils.llm_cost_calc.utils import get_web_search_requests @@ -54,7 +53,8 @@ def test_get_cost_for_anthropic_web_search_with_dict_server_tool_use(): info = _make_model_info(cost_per_query=0.01) cost = get_cost_for_anthropic_web_search( - model_info=info, usage=usage # type: ignore[arg-type] + model_info=info, + usage=usage, # type: ignore[arg-type] ) assert cost == pytest.approx(0.03) @@ -65,7 +65,8 @@ def test_get_cost_for_anthropic_web_search_with_pydantic_server_tool_use(): info = _make_model_info(cost_per_query=0.01) cost = get_cost_for_anthropic_web_search( - model_info=info, usage=usage # type: ignore[arg-type] + model_info=info, + usage=usage, # type: ignore[arg-type] ) assert cost == pytest.approx(0.03) @@ -76,7 +77,8 @@ def test_get_cost_for_anthropic_web_search_with_none_server_tool_use(): info = _make_model_info(cost_per_query=0.01) cost = get_cost_for_anthropic_web_search( - model_info=info, usage=usage # type: ignore[arg-type] + model_info=info, + usage=usage, # type: ignore[arg-type] ) assert cost == 0.0 diff --git a/tests/unit/llms/anthropic/test_count_tokens_oauth.py b/tests/unit/llms/anthropic/test_count_tokens_oauth.py index bcfc56577eb..96d909a4b3f 100644 --- a/tests/unit/llms/anthropic/test_count_tokens_oauth.py +++ b/tests/unit/llms/anthropic/test_count_tokens_oauth.py @@ -1,86 +1,271 @@ """ -Tests for Anthropic CountTokens API OAuth token handling. +Tests for the credential every Anthropic count-tokens request carries. -Verifies that get_required_headers() correctly handles OAuth tokens -(sk-ant-oat*) by delegating to optionally_handle_anthropic_oauth(). +The count-tokens handler receives the auth header that ``AnthropicModelInfo.get_auth_header`` +resolved, so a static key, an OAuth token (sk-ant-oat*), ``ANTHROPIC_AUTH_TOKEN`` and a minted +workload-identity token all reach Anthropic exactly the way chat on the same deployment does. -Regression test for https://github.com/BerriAI/litellm/issues/22040 +Regression tests for https://github.com/BerriAI/litellm/issues/22040 and for the +``ANTHROPIC_AUTH_TOKEN`` gap where count-tokens skipped minting but forwarded no credential. """ import os import sys -sys.path.insert( - 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../..")) -) +import httpx +import pytest +import respx +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../../.."))) + +import litellm +from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION +from litellm.llms.anthropic.common_utils import AnthropicModelInfo from litellm.llms.anthropic.count_tokens.transformation import ( AnthropicCountTokensConfig, ) +from litellm.types.llms.anthropic import ANTHROPIC_OAUTH_BETA_HEADER # Fake tokens for testing (not real secrets) FAKE_OAUTH_TOKEN = "sk-ant-oat01-fake-token-for-testing-123456789abcdef" FAKE_REGULAR_KEY = "sk-ant-api03-regular-key-for-testing-123456789" +FEDERATED_DEPLOYMENT = { + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "fdrl_x", + "anthropic_organization_id": "org-x", + } +} + + +def count_tokens_headers_for(api_key: str) -> dict[str, str]: + auth_header = AnthropicModelInfo.get_auth_header(api_key=api_key) + assert auth_header is not None + return AnthropicCountTokensConfig().get_count_tokens_headers(auth_header) + + +@pytest.fixture +def httpx_transport_clients(monkeypatch): + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + client_cache = getattr(litellm, "in_memory_llm_clients_cache", None) + if client_cache is not None: + client_cache.flush_cache() + yield + if client_cache is not None: + client_cache.flush_cache() + class TestCountTokensOAuthHeaders: """Tests that count_tokens headers are correct for both regular and OAuth keys.""" def test_regular_api_key_uses_x_api_key(self): """Regular API keys should be sent via x-api-key header.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_REGULAR_KEY) + headers = count_tokens_headers_for(FAKE_REGULAR_KEY) assert headers["x-api-key"] == FAKE_REGULAR_KEY assert "authorization" not in headers def test_oauth_key_uses_bearer_authorization(self): """OAuth tokens (sk-ant-oat*) should be sent via Authorization: Bearer.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_OAUTH_TOKEN) + headers = count_tokens_headers_for(FAKE_OAUTH_TOKEN) assert headers.get("authorization") == f"Bearer {FAKE_OAUTH_TOKEN}" assert "x-api-key" not in headers def test_oauth_key_sets_oauth_beta_header(self): """OAuth tokens should trigger the anthropic-beta oauth header.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_OAUTH_TOKEN) + headers = count_tokens_headers_for(FAKE_OAUTH_TOKEN) - assert "oauth-2025-04-20" in headers.get("anthropic-beta", "") + assert ANTHROPIC_OAUTH_BETA_HEADER in headers.get("anthropic-beta", "").split(",") def test_regular_key_preserves_token_counting_beta(self): """Regular keys should keep the token-counting beta header.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_REGULAR_KEY) + headers = count_tokens_headers_for(FAKE_REGULAR_KEY) - assert "token-counting" in headers.get("anthropic-beta", "") + assert headers.get("anthropic-beta") == ANTHROPIC_TOKEN_COUNTING_BETA_VERSION def test_headers_always_have_content_type(self): """Both regular and OAuth paths should have Content-Type.""" - config = AnthropicCountTokensConfig() - for key in [FAKE_REGULAR_KEY, FAKE_OAUTH_TOKEN]: - headers = config.get_required_headers(key) + headers = count_tokens_headers_for(key) assert headers["Content-Type"] == "application/json" def test_headers_always_have_anthropic_version(self): """Both paths should have anthropic-version.""" - config = AnthropicCountTokensConfig() - for key in [FAKE_REGULAR_KEY, FAKE_OAUTH_TOKEN]: - headers = config.get_required_headers(key) + headers = count_tokens_headers_for(key) assert headers["anthropic-version"] == "2023-06-01" def test_oauth_key_preserves_token_counting_beta(self): """OAuth tokens must preserve the token-counting beta alongside the OAuth beta.""" - config = AnthropicCountTokensConfig() - headers = config.get_required_headers(FAKE_OAUTH_TOKEN) + headers = count_tokens_headers_for(FAKE_OAUTH_TOKEN) - beta_value = headers.get("anthropic-beta", "") - assert ( - "token-counting" in beta_value - ), f"token-counting beta missing from OAuth headers: {beta_value}" - assert ( - "oauth-2025-04-20" in beta_value - ), f"oauth beta missing from OAuth headers: {beta_value}" + betas = headers.get("anthropic-beta", "").split(",") + assert ANTHROPIC_TOKEN_COUNTING_BETA_VERSION in betas, f"token-counting beta missing: {betas}" + assert ANTHROPIC_OAUTH_BETA_HEADER in betas, f"oauth beta missing: {betas}" + + +class TestCountTokensUsesWorkloadIdentity: + """A federated deployment holds no static key. Without minting one, count_tokens returns None + and the caller silently falls back to the local tokenizer, so the number a federated + deployment reports would never come from Anthropic.""" + + @pytest.mark.asyncio + async def test_a_federated_deployment_mints_and_counts(self, monkeypatch): + from litellm.llms.anthropic.count_tokens import token_counter as token_counter_module + + monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) + minted = "sk-ant-oat01-minted-for-count" + + async def fake_mint(_params, _api_base, _model): + return minted + + monkeypatch.setattr("litellm.llms.anthropic.common_utils.aget_anthropic_wif_token", fake_mint) + + seen: dict[str, object] = {} + + async def fake_request(**kwargs): + seen.update(kwargs) + return {"input_tokens": 42} + + monkeypatch.setattr( + token_counter_module.anthropic_count_tokens_handler, + "handle_count_tokens_request", + fake_request, + raising=False, + ) + + result = await token_counter_module.AnthropicTokenCounter().count_tokens( + model_to_use="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + contents=None, + deployment=FEDERATED_DEPLOYMENT, + request_model="claude-sonnet-4-5", + ) + + assert result is not None + assert result.total_tokens == 42 + assert seen["auth_header"] == { + "authorization": f"Bearer {minted}", + "anthropic-beta": ANTHROPIC_OAUTH_BETA_HEADER, + } + + @pytest.mark.asyncio + async def test_an_auth_token_deployment_counts_with_a_bearer_and_never_mints( + self, monkeypatch, httpx_transport_clients + ): + """With only ``ANTHROPIC_AUTH_TOKEN`` set, chat on a federated deployment authenticates with + that token, so count-tokens must send the same Bearer instead of silently returning None.""" + from litellm.llms.anthropic.count_tokens import token_counter as token_counter_module + + for name in ("ANTHROPIC_API_KEY", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"): + monkeypatch.delenv(name, raising=False) + monkeypatch.setenv("ANTHROPIC_AUTH_TOKEN", "bearer-token-for-testing") + + async def fake_mint(_params, _api_base, _model): + raise AssertionError("an auth-token deployment must never mint a federated token") + + monkeypatch.setattr("litellm.llms.anthropic.common_utils.aget_anthropic_wif_token", fake_mint) + + with respx.mock(assert_all_called=True) as router: + route = router.post("https://api.anthropic.com/v1/messages/count_tokens").mock( + return_value=httpx.Response(200, json={"input_tokens": 11}) + ) + result = await token_counter_module.AnthropicTokenCounter().count_tokens( + model_to_use="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + contents=None, + deployment=FEDERATED_DEPLOYMENT, + request_model="claude-sonnet-4-5", + ) + + assert result is not None + assert result.total_tokens == 11 + assert result.tokenizer_type == "anthropic_api" + sent = route.calls.last.request.headers + assert sent["authorization"] == "Bearer bearer-token-for-testing" + assert "x-api-key" not in sent + betas = sent["anthropic-beta"].split(",") + assert ANTHROPIC_TOKEN_COUNTING_BETA_VERSION in betas + assert ANTHROPIC_OAUTH_BETA_HEADER not in betas + + @pytest.mark.asyncio + async def test_a_failed_mint_degrades_like_an_anthropic_error(self, monkeypatch): + from litellm.llms.anthropic.count_tokens import token_counter as token_counter_module + + monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) + + async def failing_mint(_params, _api_base, model): + raise litellm.AuthenticationError( + message="federation_rule_id is not a well-formed fdrl_ tagged ID", + llm_provider="anthropic", + model=model, + ) + + monkeypatch.setattr("litellm.llms.anthropic.common_utils.aget_anthropic_wif_token", failing_mint) + + result = await token_counter_module.AnthropicTokenCounter().count_tokens( + model_to_use="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + contents=None, + deployment={ + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5", + "anthropic_federation_rule_id": "not-a-rule", + "anthropic_organization_id": "org-x", + } + }, + request_model="claude-sonnet-4-5", + ) + + assert result is not None + assert result.error is True + assert result.status_code == 401 + assert result.total_tokens == 0 + assert "fdrl_" in (result.error_message or "") + + @pytest.mark.asyncio + async def test_a_vault_backed_static_key_never_mints(self, monkeypatch): + from litellm.llms.anthropic.count_tokens import token_counter as token_counter_module + + monkeypatch.delenv("ANTHROPIC_API_KEY", raising=False) + monkeypatch.delenv("ANTHROPIC_AUTH_TOKEN", raising=False) + vault_key = "sk-ant-api03-only-in-the-vault" + + def vault_only(secret_name, default_value=None): + return vault_key if secret_name == "ANTHROPIC_API_KEY" else None + + monkeypatch.setattr("litellm.secret_managers.main.get_secret_str", vault_only, raising=False) + + async def fake_mint(_params, _api_base, _model): + raise AssertionError("a static key must never mint a federated token") + + monkeypatch.setattr("litellm.llms.anthropic.common_utils.aget_anthropic_wif_token", fake_mint) + + seen: dict[str, object] = {} + + async def fake_request(**kwargs): + seen.update(kwargs) + return {"input_tokens": 7} + + monkeypatch.setattr( + token_counter_module.anthropic_count_tokens_handler, + "handle_count_tokens_request", + fake_request, + raising=False, + ) + + result = await token_counter_module.AnthropicTokenCounter().count_tokens( + model_to_use="claude-sonnet-4-5", + messages=[{"role": "user", "content": "hi"}], + contents=None, + deployment=FEDERATED_DEPLOYMENT, + request_model="claude-sonnet-4-5", + ) + + assert result is not None + assert result.total_tokens == 7 + assert seen["auth_header"] == {"x-api-key": vault_key} diff --git a/tests/unit/llms/anthropic/test_message_sanitization.py b/tests/unit/llms/anthropic/test_message_sanitization.py index 79ed321d0ee..7afa60baf7e 100644 --- a/tests/unit/llms/anthropic/test_message_sanitization.py +++ b/tests/unit/llms/anthropic/test_message_sanitization.py @@ -12,9 +12,7 @@ import sys import os # Add the parent directory to the path so we can import litellm -sys.path.insert( - 0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../..")) -) +sys.path.insert(0, os.path.abspath(os.path.join(os.path.dirname(__file__), "../../../.."))) import litellm from litellm.litellm_core_utils.prompt_templates.factory import ( @@ -68,10 +66,7 @@ class TestMessageSanitization: assert sanitized[1]["role"] == "assistant" assert sanitized[2]["role"] == "tool" assert sanitized[2]["tool_call_id"] == "toolu_01Kus2cC3ydjBW7UK4GJqBP4" - assert ( - "skipped" in sanitized[2]["content"].lower() - or "interrupted" in sanitized[2]["content"].lower() - ) + assert "skipped" in sanitized[2]["content"].lower() or "interrupted" in sanitized[2]["content"].lower() assert "get_weather" in sanitized[2]["content"] def test_case_a_orphaned_tool_call_multiple(self): @@ -115,12 +110,8 @@ class TestMessageSanitization: assert len(sanitized) == 4 assert sanitized[0]["role"] == "user" assert sanitized[1]["role"] == "assistant" - assert ( - sanitized[2]["tool_call_id"] == "call_1" - ) # Original tool result (first in tool_calls) - assert ( - sanitized[3]["tool_call_id"] == "call_2" - ) # Dummy added for missing call_2 + assert sanitized[2]["tool_call_id"] == "call_1" # Original tool result (first in tool_calls) + assert sanitized[3]["tool_call_id"] == "call_2" # Dummy added for missing call_2 def test_case_b_orphaned_tool_result(self): """ @@ -188,10 +179,7 @@ class TestMessageSanitization: assert len(sanitized) == 2 assert sanitized[0]["role"] == "user" - assert ( - sanitized[0]["content"] - == "[System: Empty message content sanitised to satisfy protocol]" - ) + assert sanitized[0]["content"] == "[System: Empty message content sanitised to satisfy protocol]" def test_case_c_whitespace_only_content(self): """ @@ -206,14 +194,8 @@ class TestMessageSanitization: sanitized = sanitize_messages_for_tool_calling(messages) assert len(sanitized) == 2 - assert ( - sanitized[0]["content"] - == "[System: Empty message content sanitised to satisfy protocol]" - ) - assert ( - sanitized[1]["content"] - == "[System: Empty message content sanitised to satisfy protocol]" - ) + assert sanitized[0]["content"] == "[System: Empty message content sanitised to satisfy protocol]" + assert sanitized[1]["content"] == "[System: Empty message content sanitised to satisfy protocol]" def test_case_c_valid_content_preserved(self): """ @@ -270,10 +252,7 @@ class TestMessageSanitization: assert sanitized[2]["role"] == "tool" assert sanitized[2]["tool_call_id"] == "call_1" # Dummy added assert sanitized[3]["role"] == "user" - assert ( - sanitized[3]["content"] - == "[System: Empty message content sanitised to satisfy protocol]" - ) + assert sanitized[3]["content"] == "[System: Empty message content sanitised to satisfy protocol]" assert sanitized[4]["role"] == "assistant" def test_modify_params_false_no_sanitization(self): @@ -329,9 +308,7 @@ class TestMessageSanitization: ] # This should not raise an error and should add dummy tool result - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-5", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-5", llm_provider="anthropic") # Should have at least 2 messages (user and assistant) # The tool result will be merged into user content @@ -355,23 +332,17 @@ class TestMessageSanitization: {"role": "user", "content": ""}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-5", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-5", llm_provider="anthropic") # All three user messages get merged into one user turn for Anthropic. assert len(result) == 1 assert result[0]["role"] == "user" - text_blocks = [ - b for b in result[0]["content"] if isinstance(b, dict) and b.get("type") == "text" - ] + text_blocks = [b for b in result[0]["content"] if isinstance(b, dict) and b.get("type") == "text"] assert len(text_blocks) == 3 # No text block may be empty — that's the contract Anthropic enforces. for block in text_blocks: assert block["text"].strip() != "" - assert text_blocks[2]["text"] == ( - "[System: Empty message content sanitised to satisfy protocol]" - ) + assert text_blocks[2]["text"] == ("[System: Empty message content sanitised to satisfy protocol]") def test_empty_text_block_in_list_content_sanitized(self): """ @@ -392,14 +363,10 @@ class TestMessageSanitization: }, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-5", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-5", llm_provider="anthropic") assert len(result) == 1 - text_blocks = [ - b for b in result[0]["content"] if isinstance(b, dict) and b.get("type") == "text" - ] + text_blocks = [b for b in result[0]["content"] if isinstance(b, dict) and b.get("type") == "text"] assert len(text_blocks) == 3 assert text_blocks[0]["text"] == "real content" for block in text_blocks[1:]: @@ -418,9 +385,7 @@ class TestMessageSanitization: {"role": "user", "content": "How are you?"}, ] - result = anthropic_messages_pt( - messages=messages, model="claude-sonnet-4-5", llm_provider="anthropic" - ) + result = anthropic_messages_pt(messages=messages, model="claude-sonnet-4-5", llm_provider="anthropic") # Two user turns + one assistant turn (alternation preserved). assert len(result) == 3 diff --git a/tests/unit/llms/base_llm/auth/__init__.py b/tests/unit/llms/base_llm/auth/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/unit/llms/base_llm/auth/test_client_credentials.py b/tests/unit/llms/base_llm/auth/test_client_credentials.py new file mode 100644 index 00000000000..76c309cae37 --- /dev/null +++ b/tests/unit/llms/base_llm/auth/test_client_credentials.py @@ -0,0 +1,484 @@ +import base64 +import logging +from collections.abc import Mapping +from typing import Final +from urllib.parse import parse_qsl, unquote + +import httpx +import pytest + +from litellm.llms.base_llm.auth.client_credentials import ( + _HttpxSyncKeycloakPoster, + _default_secret_reader, + _new_keycloak_handler, + fetch_keycloak_assertion, + keycloak_assertion_source, +) +from litellm.llms.base_llm.auth.identity_source import KeycloakSource, identity_source_ref +from litellm.llms.base_llm.auth.token_exchange import MAX_RESPONSE_BYTES + +TOKEN_URL: Final = "https://keycloak.example/realms/litellm/protocol/openid-connect/token" +CLIENT_ID: Final = "litellm" +CLIENT_SECRET_REF: Final = "oidc/env/KEYCLOAK_CLIENT_SECRET" +CLIENT_SECRET: Final = "s3cr3t-client-value" + + +class RecordedRequest: + def __init__(self, url: str, content: bytes, headers: Mapping[str, str], timeout: float) -> None: + self.url = url + self.content = content + self.headers = dict(headers) + self.timeout = timeout + + def form_body(self) -> dict[str, str]: + return dict(parse_qsl(self.content.decode())) + + +class ScriptedPoster: + """Returns one scripted response per call; records every request it receives.""" + + def __init__(self, responses: list[httpx.Response]) -> None: + self.requests: list[RecordedRequest] = [] + self._responses = list(responses) + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.requests.append(RecordedRequest(url, content, headers, timeout)) + return self._responses.pop(0) if len(self._responses) > 1 else self._responses[0] + + +class RaisingPoster: + def __init__(self, error: Exception) -> None: + self.calls = 0 + self._error = error + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.calls += 1 + raise self._error + + +def make_config( + auth_method: str = "client_secret_basic", + scope: str | None = None, + token_url: str = TOKEN_URL, + client_secret_ref: str = CLIENT_SECRET_REF, + client_id: str = CLIENT_ID, +) -> KeycloakSource: + return KeycloakSource( + token_url=token_url, + client_id=client_id, + client_secret_ref=client_secret_ref, + auth_method=auth_method, # pyright: ignore[reportArgumentType] # test-only string widened for parametrization + scope=scope, + ) + + +def secret_reader_returning(secret: str | None): + def reader(ref: str) -> str | None: + assert ref == CLIENT_SECRET_REF + return secret + + return reader + + +DEFAULT_SECRET_READER: Final = secret_reader_returning(CLIENT_SECRET) + + +def token_response(access_token: str = "keycloak-minted-token") -> httpx.Response: + return httpx.Response(200, json={"access_token": access_token, "token_type": "Bearer", "expires_in": 300}) + + +class TestClientSecretBasic: + def test_sends_basic_auth_header_and_no_secret_in_body(self): + poster = ScriptedPoster([token_response("minted-1")]) + + token = fetch_keycloak_assertion( + make_config(auth_method="client_secret_basic"), poster=poster, secret_reader=DEFAULT_SECRET_READER + ) + + assert token == "minted-1" + request = poster.requests[0] + assert request.url == TOKEN_URL + expected_auth = "Basic " + base64.b64encode(f"{CLIENT_ID}:{CLIENT_SECRET}".encode()).decode("ascii") + assert request.headers["authorization"] == expected_auth + assert request.headers["content-type"] == "application/x-www-form-urlencoded" + body = request.form_body() + assert body["grant_type"] == "client_credentials" + assert "client_secret" not in body + assert "client_id" not in body + + def test_reserved_characters_are_form_encoded_before_basic(self): + """RFC 6749 2.3.1 requires the client id and secret be application/x-www-form-urlencoded + (Appendix B) before being base64'd into the Basic header; a raw join lets a reserved + character in either value corrupt the ':'-joined pair Keycloak decodes back out.""" + client_id = "id:with+reserved% chars" + client_secret = "secret:with+reserved% chars" + poster = ScriptedPoster([token_response("minted-reserved")]) + + fetch_keycloak_assertion( + make_config(auth_method="client_secret_basic", client_id=client_id), + poster=poster, + secret_reader=secret_reader_returning(client_secret), + ) + + header = poster.requests[0].headers["authorization"] + assert header.startswith("Basic ") + decoded = base64.b64decode(header.removeprefix("Basic ")).decode("ascii") + encoded_id, _, encoded_secret = decoded.partition(":") + assert unquote(encoded_id) == client_id + assert unquote(encoded_secret) == client_secret + + def test_scope_included_only_when_set(self): + poster = ScriptedPoster([token_response()]) + fetch_keycloak_assertion( + make_config(scope="openid profile"), poster=poster, secret_reader=DEFAULT_SECRET_READER + ) + + assert poster.requests[0].form_body()["scope"] == "openid profile" + + poster_no_scope = ScriptedPoster([token_response()]) + fetch_keycloak_assertion(make_config(scope=None), poster=poster_no_scope, secret_reader=DEFAULT_SECRET_READER) + + assert "scope" not in poster_no_scope.requests[0].form_body() + + +class TestClientSecretPost: + def test_sends_client_id_and_secret_in_body_with_no_basic_header(self): + poster = ScriptedPoster([token_response("minted-2")]) + + token = fetch_keycloak_assertion( + make_config(auth_method="client_secret_post"), poster=poster, secret_reader=DEFAULT_SECRET_READER + ) + + assert token == "minted-2" + request = poster.requests[0] + assert "authorization" not in request.headers + body = request.form_body() + assert body["grant_type"] == "client_credentials" + assert body["client_id"] == CLIENT_ID + assert body["client_secret"] == CLIENT_SECRET + + +class TestOnePostPerExchange: + def test_exactly_one_post_per_call_no_cache(self): + poster = ScriptedPoster([token_response("first"), token_response("second")]) + + first = fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + second = fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert first == "first" + assert second == "second" + assert len(poster.requests) == 2 + + +class TestInvalidClient: + def test_400_invalid_client_surfaces_redacted_detail(self): + poster = ScriptedPoster( + [httpx.Response(400, json={"error": "invalid_client", "error_description": "unauthorized client"})] + ) + + with pytest.raises(ValueError, match="invalid_client") as exc_info: + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert "unauthorized client" in str(exc_info.value) + assert "400" in str(exc_info.value) + assert CLIENT_SECRET not in str(exc_info.value) + + def test_echoed_client_secret_is_never_reflected_into_the_error(self): + """A misbehaving Keycloak that echoes the submitted client_secret back in its error body + must never leak it into the exception the caller sees.""" + long_secret: Final = "reflectable-secret-0123456789" + poster = ScriptedPoster( + [httpx.Response(400, json={"error": "invalid_client", "error_description": f"got {long_secret} in body"})] + ) + + with pytest.raises(ValueError, match="keycloak") as exc_info: + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=secret_reader_returning(long_secret)) + + assert long_secret not in str(exc_info.value) + assert "redacted" in str(exc_info.value) + + def test_echoed_short_client_secret_is_never_reflected_into_the_error(self): + """Real Keycloak client secrets are often shorter than a JWT: the reflection probe must + not silently stop protecting a secret just because it is under the probe's usual length.""" + short_secret: Final = "hand-set-14ch" + poster = ScriptedPoster( + [httpx.Response(400, json={"error": "invalid_client", "error_description": f"got {short_secret} in body"})] + ) + + with pytest.raises(ValueError, match="keycloak") as exc_info: + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=secret_reader_returning(short_secret)) + + assert short_secret not in str(exc_info.value) + assert "redacted" in str(exc_info.value) + + +class TestUnreachable: + def test_transport_failure_raises_diagnosable_value_error(self): + poster = RaisingPoster(httpx.ConnectError("connection refused")) + + with pytest.raises(ValueError, match="ConnectError") as exc_info: + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert poster.calls == 1 + assert CLIENT_SECRET not in str(exc_info.value) + + +class TestNon2xx: + def test_500_raises_value_error_with_status_code(self): + poster = ScriptedPoster([httpx.Response(500, json={"error": "server_error"})]) + + with pytest.raises(ValueError, match="500"): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + +class TestResponseValidation: + def test_missing_access_token_is_a_value_error(self): + poster = ScriptedPoster([httpx.Response(200, json={"token_type": "Bearer"})]) + + with pytest.raises(ValueError, match="schema validation"): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + def test_empty_access_token_is_a_value_error(self): + poster = ScriptedPoster([httpx.Response(200, json={"access_token": " "})]) + + with pytest.raises(ValueError, match="empty access_token"): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + +class TestInsecureTokenUrl: + def test_http_url_is_rejected_before_any_post(self): + poster = ScriptedPoster([token_response()]) + + with pytest.raises(ValueError, match="https"): + fetch_keycloak_assertion( + make_config(token_url="http://keycloak.example/token"), + poster=poster, + secret_reader=DEFAULT_SECRET_READER, + ) + + assert poster.requests == [] + + +class TestMissingClientSecret: + def test_unresolvable_secret_ref_raises_value_error_naming_the_ref_not_a_secret(self): + poster = ScriptedPoster([token_response()]) + + with pytest.raises(ValueError, match=CLIENT_SECRET_REF): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=secret_reader_returning(None)) + + assert poster.requests == [] + + +class TestKeycloakAssertionSource: + def test_returns_a_callable_that_fetches_fresh_each_call(self): + poster = ScriptedPoster([token_response("first"), token_response("second")]) + source = keycloak_assertion_source(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert source() == "first" + assert source() == "second" + assert len(poster.requests) == 2 + + def test_propagates_the_underlying_fetch_failure(self): + poster = ScriptedPoster([httpx.Response(400, json={"error": "invalid_client"})]) + source = keycloak_assertion_source(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + + with pytest.raises(ValueError, match="invalid_client"): + source() + + +class TestClientSecretNeverLeaks: + """Regression coverage for the load-bearing property: a Keycloak client_secret must never + surface in the assertion_ref, in any error message, or in a log record, however it fails.""" + + def test_never_in_the_assertion_ref(self): + config = make_config(client_secret_ref=CLIENT_SECRET_REF) + + ref = identity_source_ref(config) + + assert CLIENT_SECRET not in ref + assert CLIENT_SECRET_REF not in ref + + def test_never_in_any_raised_error_message_across_every_failure_mode(self): + config = make_config() + failures = [ + lambda: fetch_keycloak_assertion( + config, + poster=ScriptedPoster([httpx.Response(400, json={"error": "invalid_client"})]), + secret_reader=DEFAULT_SECRET_READER, + ), + lambda: fetch_keycloak_assertion( + config, poster=RaisingPoster(httpx.ConnectError("boom")), secret_reader=DEFAULT_SECRET_READER + ), + lambda: fetch_keycloak_assertion( + config, + poster=ScriptedPoster([httpx.Response(500, json={"error": "server_error"})]), + secret_reader=DEFAULT_SECRET_READER, + ), + lambda: fetch_keycloak_assertion( + config, poster=ScriptedPoster([token_response()]), secret_reader=secret_reader_returning(None) + ), + ] + for fail in failures: + with pytest.raises(ValueError, match="keycloak") as exc_info: + fail() + assert CLIENT_SECRET not in str(exc_info.value) + + def test_never_in_a_log_record(self, caplog: pytest.LogCaptureFixture): + with caplog.at_level(logging.DEBUG): + poster = ScriptedPoster( + [httpx.Response(400, json={"error": "invalid_client", "error_description": CLIENT_SECRET})] + ) + with pytest.raises(ValueError, match="keycloak"): + fetch_keycloak_assertion(make_config(), poster=poster, secret_reader=DEFAULT_SECRET_READER) + fetch_keycloak_assertion( + make_config(), poster=ScriptedPoster([token_response()]), secret_reader=DEFAULT_SECRET_READER + ) + + assert CLIENT_SECRET not in caplog.text + + +class StubHandler: + """Stands in for the HTTPHandler the default poster builds, so the poster's own contract + (redirects off, error responses returned rather than raised, no-response guarded) is testable + without a socket.""" + + def __init__(self, result: httpx.Response | Exception | None) -> None: + self.calls: list[dict[str, object]] = [] + self._result = result + + def post(self, url: str, *, content: bytes, headers: dict[str, str], timeout: float) -> httpx.Response | None: + self.calls.append({"url": url, "content": content, "headers": headers, "timeout": timeout}) + if isinstance(self._result, Exception): + raise self._result + return self._result + + +class TestDefaultKeycloakPoster: + def test_builds_its_handler_once_with_redirects_disabled(self): + built: list[StubHandler] = [] + + def factory() -> StubHandler: + handler = StubHandler(httpx.Response(200, json={"access_token": "kc-token"})) + built.append(handler) + return handler + + poster: Final = _HttpxSyncKeycloakPoster(handler_factory=factory) # pyright: ignore[reportArgumentType] # StubHandler stands in for the legacy-untyped HTTPHandler + for _ in range(3): + poster.post(TOKEN_URL, content=b"grant_type=client_credentials", headers={}, timeout=1.0) + + assert len(built) == 1, "the handler is built once and reused" + assert len(built[0].calls) == 3 + + def test_the_real_handler_refuses_to_follow_redirects(self): + handler: Final = _new_keycloak_handler() + assert handler.client.follow_redirects is False, ( + "a redirected token POST would replay the client secret to whatever host the redirect names" + ) + + def test_an_http_status_error_becomes_its_response_rather_than_an_exception(self): + response: Final = httpx.Response( + 401, json={"error": "invalid_client"}, request=httpx.Request("POST", TOKEN_URL) + ) + poster: Final = _HttpxSyncKeycloakPoster( + handler_factory=lambda: StubHandler( + httpx.HTTPStatusError("boom", request=response.request, response=response) + ) # pyright: ignore[reportArgumentType] # StubHandler stands in for the legacy-untyped HTTPHandler + ) + + assert poster.post(TOKEN_URL, content=b"", headers={}, timeout=1.0).status_code == 401 + + def test_a_missing_response_is_a_transport_error_not_a_none_deref(self): + poster: Final = _HttpxSyncKeycloakPoster(handler_factory=lambda: StubHandler(None)) # pyright: ignore[reportArgumentType] # StubHandler stands in for the legacy-untyped HTTPHandler + + with pytest.raises(httpx.TransportError): + poster.post(TOKEN_URL, content=b"", headers={}, timeout=1.0) + + +class TestDefaultSecretReader: + def test_reads_through_litellm_secret_resolution(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("KEYCLOAK_CLIENT_SECRET_FOR_DEFAULT_READER", CLIENT_SECRET) + + assert _default_secret_reader("os.environ/KEYCLOAK_CLIENT_SECRET_FOR_DEFAULT_READER") == CLIENT_SECRET + + def test_an_unset_reference_reads_as_none_so_the_caller_raises(self): + assert _default_secret_reader("os.environ/DEFINITELY_NOT_SET_KEYCLOAK_SECRET_REF") is None + + +class TestOversizedSuccessBody: + def test_a_success_body_over_the_cap_is_refused_before_it_is_parsed(self): + oversized: Final = httpx.Response(200, content=b'{"access_token": "' + b"x" * MAX_RESPONSE_BYTES + b'"}') + + with pytest.raises(ValueError, match="exceeded the size cap"): + fetch_keycloak_assertion( + make_config(), poster=ScriptedPoster([oversized]), secret_reader=DEFAULT_SECRET_READER + ) + + +class TestUnresolvedSecretRefIsNotEchoed: + """An operator who pastes the secret itself into the *_ref field turns that field INTO the + secret, and this error reaches model callers, so it must never echo the value.""" + + def test_keycloak_ref_value_is_not_in_the_error(self): + from litellm.llms.base_llm.auth.client_credentials import keycloak_assertion_source + from litellm.llms.base_llm.auth.identity_source import KeycloakSource + + pasted_secret = "sUp3r-s3cret-value-not-a-pointer" + config = KeycloakSource( + token_url="https://keycloak.example.com/realms/p/protocol/openid-connect/token", + client_id="litellm", + client_secret_ref=pasted_secret, + ) + + with pytest.raises(ValueError, match="could not be read") as excinfo: + keycloak_assertion_source(config, secret_reader=lambda _ref: None)() + + assert pasted_secret not in str(excinfo.value) + assert "withheld" in str(excinfo.value) + + def test_internal_issuer_ref_value_is_not_in_the_error(self): + from litellm.llms.base_llm.auth.identity_source import InternalIssuerSource + from litellm.llms.base_llm.auth.internal_issuer import internal_issuer_assertion_source + + pasted_pem = "-----BEGIN PRIVATE KEY-----MIGHAgEA-----END PRIVATE KEY-----" + config = InternalIssuerSource( + issuer_url="https://proxy.example.com", + subject="litellm-proxy", + signing_key_ref=pasted_pem, + ) + + with pytest.raises(ValueError, match="could not be read") as excinfo: + internal_issuer_assertion_source(config, key_reader=lambda _ref: None)() + + assert pasted_pem not in str(excinfo.value) + assert "withheld" in str(excinfo.value) + + +class TestTokenUrlIsNotEchoedWholesale: + """A token endpoint is configuration and naming it makes the error actionable, but nothing + stops an operator putting a credential in the URL, and these errors reach model callers.""" + + def test_query_string_is_dropped_from_a_status_error(self): + from litellm.llms.base_llm.auth.token_exchange import endpoint_url_for_error_message + + rendered = endpoint_url_for_error_message("https://idp.example/token?client_secret=supersecret") + + assert "supersecret" not in rendered + assert rendered == "https://idp.example/token" + + def test_userinfo_is_dropped_too(self): + from litellm.llms.base_llm.auth.token_exchange import endpoint_url_for_error_message + + rendered = endpoint_url_for_error_message("https://user:pw@idp.example:8443/token") + + assert "pw" not in rendered + assert rendered == "https://idp.example:8443/token" + + def test_transport_failure_message_carries_no_query_secret(self): + poster = RaisingPoster(httpx.ConnectTimeout("timed out")) + config = make_config(token_url="https://idp.example/token?client_secret=supersecret") + + with pytest.raises(ValueError, match="could not reach the keycloak token endpoint") as excinfo: + fetch_keycloak_assertion(config, poster=poster, secret_reader=DEFAULT_SECRET_READER) + + assert "supersecret" not in str(excinfo.value) + assert "idp.example/token" in str(excinfo.value) diff --git a/tests/unit/llms/base_llm/auth/test_identity_source.py b/tests/unit/llms/base_llm/auth/test_identity_source.py new file mode 100644 index 00000000000..bfa8847c69a --- /dev/null +++ b/tests/unit/llms/base_llm/auth/test_identity_source.py @@ -0,0 +1,239 @@ +from types import MappingProxyType +from typing import Final, Literal + +import pytest +from pydantic import ValidationError + +from litellm.llms.base_llm.auth.identity_source import ( + AnthropicIdentitySourceKind, + InternalIssuerSource, + KeycloakSource, + identity_source_config_adapter, + identity_source_ref, +) + +SIGNING_KEY_REF: Final = "oidc/env/ISSUER_SIGNING_KEY_PEM" +OTHER_SIGNING_KEY_REF: Final = "oidc/env/OTHER_SIGNING_KEY_PEM" +CLIENT_SECRET_REF: Final = "oidc/env/KEYCLOAK_CLIENT_SECRET" +ISSUER_URL: Final = "https://issuer.internal.example" +SUBJECT: Final = "workload-a" +TOKEN_URL: Final = "https://keycloak.example/realms/litellm/protocol/openid-connect/token" +CLIENT_ID: Final = "litellm" + + +def make_issuer( + issuer_url: str = ISSUER_URL, + subject: str = SUBJECT, + signing_key_ref: str = SIGNING_KEY_REF, + ttl_seconds: int = 300, +) -> InternalIssuerSource: + return InternalIssuerSource( + issuer_url=issuer_url, subject=subject, signing_key_ref=signing_key_ref, ttl_seconds=ttl_seconds + ) + + +def make_keycloak( + token_url: str = TOKEN_URL, + client_id: str = CLIENT_ID, + client_secret_ref: str = CLIENT_SECRET_REF, + auth_method: Literal["client_secret_basic", "client_secret_post"] = "client_secret_basic", + scope: str | None = None, +) -> KeycloakSource: + return KeycloakSource( + token_url=token_url, + client_id=client_id, + client_secret_ref=client_secret_ref, + auth_method=auth_method, + scope=scope, + ) + + +class TestIdentitySourceRefHashing: + def test_identical_config_hashes_idempotently(self): + assert identity_source_ref(make_issuer()) == identity_source_ref(make_issuer()) + + def test_ref_is_prefixed_by_kind(self): + assert identity_source_ref(make_issuer()).startswith("oidc/internal_issuer/") + assert identity_source_ref(make_keycloak()).startswith("oidc/keycloak/") + + def test_pointer_name_change_changes_ref(self): + """Two configs differing only in which secret a pointer names must never collide, since a + stale ref would let the token exchange's outer cache key alias two different credentials.""" + first: Final = identity_source_ref(make_issuer(signing_key_ref=SIGNING_KEY_REF)) + second: Final = identity_source_ref(make_issuer(signing_key_ref=OTHER_SIGNING_KEY_REF)) + + assert first != second + + def test_non_pointer_field_change_changes_ref(self): + first: Final = identity_source_ref(make_keycloak(scope="openid")) + second: Final = identity_source_ref(make_keycloak(scope="openid profile")) + + assert first != second + + def test_ref_never_contains_the_pointer_field_values(self): + """The ref is a fixed-width hash, not a serialization of the config, so no field value - + pointer name or otherwise - can leak into the secret-free string echoed into errors.""" + ref: Final = identity_source_ref(make_issuer()) + + assert SIGNING_KEY_REF not in ref + assert "issuer.internal.example" not in ref + + def test_different_kinds_with_disjoint_fields_never_collide(self): + assert identity_source_ref(make_issuer()) != identity_source_ref(make_keycloak()) + + +class TestInternalIssuerSourceValidation: + def test_defaults(self): + source: Final = make_issuer() + + assert source.kind == AnthropicIdentitySourceKind.internal_issuer + assert source.ttl_seconds == 300 + assert source.audience is None + + def test_ttl_seconds_over_one_hour_is_rejected(self): + with pytest.raises(ValidationError): + make_issuer(ttl_seconds=3601) + + def test_ttl_seconds_at_one_hour_is_accepted(self): + assert make_issuer(ttl_seconds=3600).ttl_seconds == 3600 + + def test_non_positive_ttl_seconds_is_rejected(self): + with pytest.raises(ValidationError): + make_issuer(ttl_seconds=0) + + def test_missing_signing_key_ref_is_rejected(self): + missing_field: Final = MappingProxyType({"issuer_url": ISSUER_URL, "subject": SUBJECT}) + + with pytest.raises(ValidationError): + InternalIssuerSource.model_validate(missing_field) + + def test_keycloak_only_field_is_rejected_as_extra(self): + mixed_variant: Final = MappingProxyType( + { + "issuer_url": ISSUER_URL, + "subject": SUBJECT, + "signing_key_ref": SIGNING_KEY_REF, + "client_secret_ref": CLIENT_SECRET_REF, + } + ) + + with pytest.raises(ValidationError): + InternalIssuerSource.model_validate(mixed_variant) + + def test_is_frozen(self): + source: Final = make_issuer() + + with pytest.raises(ValidationError): + source.subject = "workload-b" + + def test_secret_pasted_into_wrong_typed_field_is_not_echoed_in_the_error(self): + """hide_input_in_errors keeps a value the operator pasted into a mistyped field out of the + validation error, so a client_secret headed for the wrong field isn't logged in the raise.""" + leaked_secret: Final = "shh-do-not-log-me" + wrong_type: Final = MappingProxyType( + { + "issuer_url": ISSUER_URL, + "subject": SUBJECT, + "signing_key_ref": SIGNING_KEY_REF, + "ttl_seconds": leaked_secret, + } + ) + + with pytest.raises(ValidationError) as exc_info: + InternalIssuerSource.model_validate(wrong_type) + + assert leaked_secret not in str(exc_info.value) + + +class TestKeycloakSourceValidation: + def test_defaults(self): + source: Final = make_keycloak() + + assert source.kind == AnthropicIdentitySourceKind.keycloak + assert source.auth_method == "client_secret_basic" + assert source.scope is None + + def test_client_secret_post_is_accepted(self): + assert make_keycloak(auth_method="client_secret_post").auth_method == "client_secret_post" + + def test_private_key_jwt_is_not_a_supported_auth_method_yet(self): + unshipped_auth_method: Final = MappingProxyType( + { + "token_url": TOKEN_URL, + "client_id": CLIENT_ID, + "client_secret_ref": CLIENT_SECRET_REF, + "auth_method": "private_key_jwt", + } + ) + + with pytest.raises(ValidationError): + KeycloakSource.model_validate(unshipped_auth_method) + + def test_audience_field_was_dropped(self): + dropped_field: Final = MappingProxyType( + { + "token_url": TOKEN_URL, + "client_id": CLIENT_ID, + "client_secret_ref": CLIENT_SECRET_REF, + "audience": "https://anthropic.example", + } + ) + + with pytest.raises(ValidationError): + KeycloakSource.model_validate(dropped_field) + + def test_missing_client_secret_ref_is_rejected(self): + missing_field: Final = MappingProxyType({"token_url": TOKEN_URL, "client_id": CLIENT_ID}) + + with pytest.raises(ValidationError): + KeycloakSource.model_validate(missing_field) + + +class TestDiscriminatedUnionParsing: + def test_parses_internal_issuer_variant(self): + parsed: Final = identity_source_config_adapter.validate_python( + MappingProxyType( + { + "kind": "internal_issuer", + "issuer_url": ISSUER_URL, + "subject": SUBJECT, + "signing_key_ref": SIGNING_KEY_REF, + } + ) + ) + + assert isinstance(parsed, InternalIssuerSource) + + def test_parses_keycloak_variant(self): + parsed: Final = identity_source_config_adapter.validate_python( + MappingProxyType( + { + "kind": "keycloak", + "token_url": TOKEN_URL, + "client_id": CLIENT_ID, + "client_secret_ref": CLIENT_SECRET_REF, + } + ) + ) + + assert isinstance(parsed, KeycloakSource) + + def test_unknown_kind_is_a_hard_error(self): + with pytest.raises(ValidationError): + identity_source_config_adapter.validate_python(MappingProxyType({"kind": "token_file"})) + + def test_mixed_variant_fields_are_a_hard_error(self): + """A keycloak field on an internal_issuer-tagged payload must fail closed rather than be + silently dropped or silently accepted as if it selected the other variant.""" + with pytest.raises(ValidationError): + identity_source_config_adapter.validate_python( + MappingProxyType( + { + "kind": "internal_issuer", + "issuer_url": ISSUER_URL, + "subject": SUBJECT, + "signing_key_ref": SIGNING_KEY_REF, + "client_secret_ref": CLIENT_SECRET_REF, + } + ) + ) diff --git a/tests/unit/llms/base_llm/auth/test_internal_issuer.py b/tests/unit/llms/base_llm/auth/test_internal_issuer.py new file mode 100644 index 00000000000..d965a356426 --- /dev/null +++ b/tests/unit/llms/base_llm/auth/test_internal_issuer.py @@ -0,0 +1,188 @@ +import json +from typing import Final + +import jwt +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec + +from litellm.llms.base_llm.auth.identity_source import InternalIssuerSource +from litellm.llms.base_llm.auth.internal_issuer import ( + internal_issuer_assertion_source, + internal_issuer_jwks_document, + mint_internal_issuer_assertion, +) +from litellm.llms.base_llm.auth.jwt_signing import build_jwks, rfc7638_thumbprint + +SIGNING_KEY_REF: Final = "oidc/env/ISSUER_SIGNING_KEY_PEM" +ISSUER_URL: Final = "https://issuer.internal.example" +SUBJECT: Final = "workload-a" + + +_PRIVATE_VALUE: Final = 90123456789012345678901234567890123456789012345678901234567890 + + +def signing_key() -> ec.EllipticCurvePrivateKey: + return ec.derive_private_key(_PRIVATE_VALUE, ec.SECP256R1()) + + +def pem_of(key: ec.EllipticCurvePrivateKey) -> str: + return key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ).decode() + + +def make_config( + issuer_url: str = ISSUER_URL, + subject: str = SUBJECT, + audience: str | None = None, + ttl_seconds: int = 300, + signing_key_ref: str = SIGNING_KEY_REF, +) -> InternalIssuerSource: + return InternalIssuerSource( + issuer_url=issuer_url, + subject=subject, + audience=audience, + ttl_seconds=ttl_seconds, + signing_key_ref=signing_key_ref, + ) + + +def key_reader_returning(pem: str | None): + def reader(ref: str) -> str | None: + assert ref == SIGNING_KEY_REF + return pem + + return reader + + +class FakeClock: + def __init__(self, value: float) -> None: + self._value: Final = value + + def __call__(self) -> float: + return self._value + + +def decode_ignoring_wall_clock(token: str, public_key: ec.EllipticCurvePublicKey) -> dict: + """Tests mint with a fixed past ``FakeClock`` and no expected audience, so PyJWT's + real-wall-clock ``exp``/``aud`` checks (irrelevant to what these tests verify) are disabled.""" + return jwt.decode(token, public_key, algorithms=["ES256"], options={"verify_exp": False, "verify_aud": False}) + + +class TestMintInternalIssuerAssertion: + def test_required_claims_and_asymmetric_alg(self): + key: Final = signing_key() + config: Final = make_config(ttl_seconds=300) + + token: Final = mint_internal_issuer_assertion( + config, key_reader=key_reader_returning(pem_of(key)), clock=FakeClock(1_700_000_000.0) + ) + header: Final = jwt.get_unverified_header(token) + claims: Final = decode_ignoring_wall_clock(token, key.public_key()) + + assert header["alg"] == "ES256" + assert claims["sub"] == SUBJECT + assert claims["iss"] == ISSUER_URL + assert claims["iat"] == 1_700_000_000 + assert claims["exp"] == 1_700_000_300 + + def test_kid_matches_the_published_jwks(self): + key: Final = signing_key() + config: Final = make_config() + + token: Final = mint_internal_issuer_assertion( + config, key_reader=key_reader_returning(pem_of(key)), clock=FakeClock(1_700_000_000.0) + ) + + header_kid: Final = jwt.get_unverified_header(token)["kid"] + published_kid: Final = build_jwks(key.public_key())["keys"][0]["kid"] + assert header_kid == published_kid == rfc7638_thumbprint(key.public_key()) + + def test_ttl_bounds_exp_minus_iat(self): + key: Final = signing_key() + config: Final = make_config(ttl_seconds=120) + + token: Final = mint_internal_issuer_assertion( + config, key_reader=key_reader_returning(pem_of(key)), clock=FakeClock(1_700_000_000.0) + ) + claims: Final = decode_ignoring_wall_clock(token, key.public_key()) + + assert claims["exp"] - claims["iat"] == 120 + + def test_audience_included_only_when_set(self): + key: Final = signing_key() + without_audience: Final = mint_internal_issuer_assertion( + make_config(audience=None), key_reader=key_reader_returning(pem_of(key)), clock=FakeClock(1_700_000_000.0) + ) + with_audience: Final = mint_internal_issuer_assertion( + make_config(audience="urn:anthropic:federation"), + key_reader=key_reader_returning(pem_of(key)), + clock=FakeClock(1_700_000_000.0), + ) + + claims_without: Final = decode_ignoring_wall_clock(without_audience, key.public_key()) + claims_with: Final = decode_ignoring_wall_clock(with_audience, key.public_key()) + assert "aud" not in claims_without + assert claims_with["aud"] == "urn:anthropic:federation" + + def test_jti_is_present_and_fresh_on_every_mint(self): + key: Final = signing_key() + config: Final = make_config() + reader: Final = key_reader_returning(pem_of(key)) + + first: Final = decode_ignoring_wall_clock( + mint_internal_issuer_assertion(config, key_reader=reader, clock=FakeClock(1_700_000_000.0)), + key.public_key(), + ) + second: Final = decode_ignoring_wall_clock( + mint_internal_issuer_assertion(config, key_reader=reader, clock=FakeClock(1_700_000_000.0)), + key.public_key(), + ) + + assert first["jti"] and second["jti"] + assert first["jti"] != second["jti"] + + def test_missing_signing_key_raises_value_error_naming_the_ref_not_a_secret(self): + with pytest.raises(ValueError, match=SIGNING_KEY_REF): + mint_internal_issuer_assertion(make_config(), key_reader=key_reader_returning(None)) + + def test_malformed_signing_key_raises_value_error(self): + with pytest.raises(ValueError, match="not a valid unencrypted PEM"): + mint_internal_issuer_assertion(make_config(), key_reader=key_reader_returning("not-a-pem")) + + +class TestInternalIssuerAssertionSource: + def test_returns_a_callable_that_mints_fresh_each_call(self): + key: Final = signing_key() + source: Final = internal_issuer_assertion_source(make_config(), key_reader=key_reader_returning(pem_of(key))) + + first: Final = jwt.decode(source(), key.public_key(), algorithms=["ES256"]) + second: Final = jwt.decode(source(), key.public_key(), algorithms=["ES256"]) + + assert first["jti"] != second["jti"] + + def test_propagates_the_underlying_mint_failure(self): + source: Final = internal_issuer_assertion_source(make_config(), key_reader=key_reader_returning(None)) + + with pytest.raises(ValueError, match=SIGNING_KEY_REF): + source() + + +class TestInternalIssuerJwksDocument: + def test_matches_the_key_used_to_mint(self): + key: Final = signing_key() + config: Final = make_config() + reader: Final = key_reader_returning(pem_of(key)) + + document: Final = json.loads(internal_issuer_jwks_document(config, key_reader=reader)) + token: Final = mint_internal_issuer_assertion(config, key_reader=reader, clock=FakeClock(1_700_000_000.0)) + + assert document["keys"][0]["kid"] == jwt.get_unverified_header(token)["kid"] + assert decode_ignoring_wall_clock(token, key.public_key()) + + def test_missing_signing_key_raises_value_error(self): + with pytest.raises(ValueError, match=SIGNING_KEY_REF): + internal_issuer_jwks_document(make_config(), key_reader=key_reader_returning(None)) diff --git a/tests/unit/llms/base_llm/auth/test_jwt_signing.py b/tests/unit/llms/base_llm/auth/test_jwt_signing.py new file mode 100644 index 00000000000..f1fbc698c82 --- /dev/null +++ b/tests/unit/llms/base_llm/auth/test_jwt_signing.py @@ -0,0 +1,215 @@ +import base64 +import hashlib +import json +import subprocess +import sys +import textwrap +import time +from typing import Final + +import jwt +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec, rsa + +from litellm.llms.base_llm.auth.jwt_signing import ( + MISSING_SIGNING_DEPENDENCIES_MESSAGE, + build_jwk, + build_jwks, + jwks_document_json, + load_es256_private_key, + rfc7638_thumbprint, + sign_es256_jwt, +) + +_FIXED_PRIVATE_VALUE: Final = 55090612345678901234567890123456789012345678901234567890123456 +_OTHER_PRIVATE_VALUE: Final = 1 + + +def fixed_private_key(value: int = _FIXED_PRIVATE_VALUE) -> ec.EllipticCurvePrivateKey: + return ec.derive_private_key(value, ec.SECP256R1()) + + +def pem_of(key: ec.EllipticCurvePrivateKey) -> str: + return key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ).decode() + + +def independent_thumbprint(public_key: ec.EllipticCurvePublicKey) -> str: + """Recomputes RFC 7638 by hand, deliberately not sharing a single line of code with + ``jwt_signing.rfc7638_thumbprint`` -- a mutation that broke the real implementation must not + also break this reference, or the two would trivially agree by sharing the bug.""" + numbers: Final = public_key.public_numbers() + x: Final = base64.urlsafe_b64encode(numbers.x.to_bytes(32, "big")).rstrip(b"=").decode() + y: Final = base64.urlsafe_b64encode(numbers.y.to_bytes(32, "big")).rstrip(b"=").decode() + canonical: Final = f'{{"crv":"P-256","kty":"EC","x":"{x}","y":"{y}"}}' + return base64.urlsafe_b64encode(hashlib.sha256(canonical.encode()).digest()).rstrip(b"=").decode() + + +class TestLoadEs256PrivateKey: + def test_valid_ec_p256_pem_loads(self): + key: Final = load_es256_private_key(pem_of(fixed_private_key())) + + assert isinstance(key, ec.EllipticCurvePrivateKey) + assert isinstance(key.curve, ec.SECP256R1) + + def test_garbage_pem_is_rejected(self): + with pytest.raises(ValueError, match="not a valid unencrypted PEM"): + load_es256_private_key("not a pem") + + def test_rsa_key_is_rejected(self): + rsa_pem: Final = ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + .decode() + ) + + with pytest.raises(ValueError, match="P-256"): + load_es256_private_key(rsa_pem) + + def test_non_p256_curve_is_rejected(self): + secp384_pem: Final = ( + ec.generate_private_key(ec.SECP384R1()) + .private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ) + .decode() + ) + + with pytest.raises(ValueError, match="P-256"): + load_es256_private_key(secp384_pem) + + def test_error_never_echoes_key_material(self): + pem: Final = pem_of(fixed_private_key()) + + with pytest.raises(ValueError, match="P-256"): + load_es256_private_key(pem_of(ec.generate_private_key(ec.SECP384R1()))) + with pytest.raises(ValueError, match="not a valid unencrypted PEM") as exc_info: + load_es256_private_key("garbage-not-a-pem") + + assert pem not in str(exc_info.value) + + +class TestRfc7638Thumbprint: + def test_matches_independent_recomputation(self): + public_key: Final = fixed_private_key().public_key() + + assert rfc7638_thumbprint(public_key) == independent_thumbprint(public_key) + + def test_different_keys_have_different_thumbprints(self): + first: Final = fixed_private_key(_FIXED_PRIVATE_VALUE).public_key() + second: Final = fixed_private_key(_OTHER_PRIVATE_VALUE).public_key() + + assert rfc7638_thumbprint(first) != rfc7638_thumbprint(second) + + def test_thumbprint_is_deterministic(self): + public_key: Final = fixed_private_key().public_key() + + assert rfc7638_thumbprint(public_key) == rfc7638_thumbprint(public_key) + + +class TestBuildJwks: + def test_jwks_contains_one_key_matching_the_thumbprint(self): + public_key: Final = fixed_private_key().public_key() + + jwks: Final = build_jwks(public_key) + + assert len(jwks["keys"]) == 1 + assert jwks["keys"][0]["kid"] == rfc7638_thumbprint(public_key) + assert jwks["keys"][0]["kty"] == "EC" + assert jwks["keys"][0]["crv"] == "P-256" + assert jwks["keys"][0]["alg"] == "ES256" + + def test_build_jwk_stamps_the_given_kid_verbatim(self): + jwk: Final = build_jwk(fixed_private_key().public_key(), kid="caller-supplied-kid") + + assert jwk["kid"] == "caller-supplied-kid" + + def test_jwks_document_json_round_trips_through_build_jwks(self): + key: Final = fixed_private_key() + + document: Final = json.loads(jwks_document_json(pem_of(key))) + jwks: Final = build_jwks(key.public_key()) + + assert document == {"keys": [dict(jwk) for jwk in jwks["keys"]]} + + +class TestSignEs256Jwt: + def test_minted_token_verifies_against_the_matching_public_key(self): + key: Final = fixed_private_key() + now: Final = int(time.time()) + claims: Final = {"sub": "workload-a", "iss": "https://issuer.example", "iat": now, "exp": now + 300} + + token: Final = sign_es256_jwt(pem_of(key), claims) + decoded: Final = jwt.decode(token, key.public_key(), algorithms=["ES256"]) + + assert decoded == claims + + def test_header_alg_is_es256(self): + token: Final = sign_es256_jwt(pem_of(fixed_private_key()), {"sub": "x"}) + + assert jwt.get_unverified_header(token)["alg"] == "ES256" + + def test_header_kid_matches_the_published_jwks(self): + key: Final = fixed_private_key() + + token: Final = sign_es256_jwt(pem_of(key), {"sub": "x"}) + + header_kid: Final = jwt.get_unverified_header(token)["kid"] + published_kid: Final = build_jwks(key.public_key())["keys"][0]["kid"] + assert header_kid == published_kid == rfc7638_thumbprint(key.public_key()) + + def test_wrong_key_fails_verification(self): + signing_key: Final = fixed_private_key(_FIXED_PRIVATE_VALUE) + other_key: Final = fixed_private_key(_OTHER_PRIVATE_VALUE) + + token: Final = sign_es256_jwt(pem_of(signing_key), {"sub": "x"}) + + with pytest.raises(jwt.exceptions.InvalidSignatureError): + jwt.decode(token, other_key.public_key(), algorithms=["ES256"]) + + +class TestBaseSdkImport: + """A base ``pip install litellm`` has neither PyJWT nor cryptography (both are proxy extras), + and ``litellm/__init__`` reaches this module through the Anthropic provider, so a + module-level import of either would break ``import litellm`` for every base SDK user.""" + + def test_module_imports_with_pyjwt_and_cryptography_absent(self): + script: Final = textwrap.dedent( + """ + import sys + + class Blocker: + def find_spec(self, name, path=None, target=None): + if name.split(".")[0] in {"jwt", "cryptography"}: + raise ModuleNotFoundError(f"No module named {name!r}") + + sys.meta_path.insert(0, Blocker()) + import litellm + from litellm.llms.base_llm.auth.jwt_signing import jwks_document_json + try: + jwks_document_json("not a key") + except ImportError as e: + print(e) + """ + ) + result: Final = subprocess.run( + [sys.executable, "-I", "-c", script], capture_output=True, text=True, check=False + ) + assert result.returncode == 0, result.stderr[-2000:] + assert result.stdout.strip() == MISSING_SIGNING_DEPENDENCIES_MESSAGE + + def test_signing_reports_the_missing_extra(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setitem(sys.modules, "jwt", None) + with pytest.raises(ImportError, match="litellm\\[proxy\\]"): + sign_es256_jwt(pem_of(fixed_private_key()), {"sub": "x"}) + diff --git a/tests/unit/llms/base_llm/auth/test_shared_token_store.py b/tests/unit/llms/base_llm/auth/test_shared_token_store.py new file mode 100644 index 00000000000..298e61c4c79 --- /dev/null +++ b/tests/unit/llms/base_llm/auth/test_shared_token_store.py @@ -0,0 +1,321 @@ +"""Two engines standing in for two uvicorn workers that read the same assertion and share one +``FileTokenStore``: an issuer that accepts each assertion once must see one exchange per assertion.""" + +import errno +import json +import os +import stat +import threading +from collections.abc import Callable, Mapping +from pathlib import Path +from typing import Final + +import httpx +import pytest + +from pydantic import SecretStr + +from litellm.llms.base_llm.auth.shared_token_store import ( + CACHE_DIR_ENV, + FileTokenStore, + StoredToken, + default_shared_token_store, +) +from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine +from litellm.llms.base_llm.auth.types import MintedToken, TokenEndpointError +from tests.unit.llms.base_llm.auth.test_token_exchange import ( + DEFAULT_ASSERTION, + DEFAULT_REF, + FakeClock, + ManualExecutor, + RecordingMetricsSink, + ScriptedPoster, + make_spec, + token_response, +) + +real_write_bytes: Final = Path.write_bytes + + +class SingleUsePoster: + """Mints for an assertion it has never seen and answers 401 to any assertion sent a second time, + which is how an issuer enforcing single-use ``jti`` behaves.""" + + def __init__(self, token: str = "sk-ant-oat01-minted", expires_in: int = 3600) -> None: + self.requests: list[dict] = [] + self._token = token + self._expires_in = expires_in + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + body = json.loads(content) + seen_before = any(prior["assertion"] == body["assertion"] for prior in self.requests) + self.requests.append(body) + if seen_before: + return httpx.Response(401, json={"error": "invalid_grant"}) + return token_response(f"{self._token}-{len(self.requests)}", expires_in=self._expires_in) + + +def store_engine( + poster, + store: FileTokenStore, + *, + reader: Mapping[str, str] | None = None, + clock: FakeClock | None = None, + wall_clock: Callable[[], float] | None = None, +) -> JwtBearerTokenExchangeEngine: + return JwtBearerTokenExchangeEngine( + poster=poster, + assertion_reader=(reader if reader is not None else {DEFAULT_REF: DEFAULT_ASSERTION}).get, + clock=clock if clock is not None else FakeClock(), + refresh_executor=ManualExecutor(), + metrics_sink=RecordingMetricsSink(), + shared_store=store, + wall_clock=wall_clock if wall_clock is not None else FakeClock(1_700_000_000.0), + ) + + +def minted(result: object) -> MintedToken: + assert isinstance(result, MintedToken), result + return result + + +def stored_files(directory: Path) -> list[Path]: + return sorted(directory.glob("*.json")) + + +def test_second_worker_reuses_the_first_workers_token_without_a_post(tmp_path: Path): + poster = SingleUsePoster() + store = FileTokenStore(tmp_path) + first = minted(store_engine(poster, store).get_token(make_spec())) + + second = minted(store_engine(poster, store).get_token(make_spec())) + + assert second.access_token.get_secret_value() == first.access_token.get_secret_value() + assert len(poster.requests) == 1 + + +def test_a_minted_assertion_never_reaches_the_shared_store(tmp_path: Path): + """internal_issuer and keycloak mint a fresh assertion per exchange, so no other worker ever holds + the same one and a stored token could never be matched back. Writing a live token to disk for a + lookup that cannot succeed is exposure that buys nothing.""" + poster = SingleUsePoster() + store = FileTokenStore(tmp_path) + assertions = iter(("minted-jwt-1", "minted-jwt-2")) + spec = make_spec(assertion_source=lambda: next(assertions)) + + first = minted(store_engine(poster, store).get_token(spec)) + second = minted(store_engine(poster, store).get_token(spec)) + + assert stored_files(tmp_path) == [], "a per-exchange assertion must keep its token off disk" + assert [request["assertion"] for request in poster.requests] == ["minted-jwt-1", "minted-jwt-2"] + assert second.access_token.get_secret_value() != first.access_token.get_secret_value() + + +def test_a_failed_write_leaves_no_staging_file_holding_a_live_token(tmp_path: Path): + """Nothing ever sweeps this directory, so a staging file a failed write leaves behind would keep a + working token readable on disk for as long as the pod lives.""" + store = FileTokenStore(tmp_path) + (tmp_path / "occupied.json").mkdir() + + store.save( + "occupied", + StoredToken(access_token=SecretStr("sk-ant-oat01-live"), expires_at_epoch=None, assertion_sha256="sha"), + ) + + leaked = [path for path in tmp_path.rglob("*") if path.is_file() and "sk-ant-oat01-live" in path.read_text()] + assert leaked == [], "a staged token file survived the failed write" + assert store.load("occupied") is None + + +def test_a_write_that_only_fails_on_close_leaves_no_staging_file(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + """A token is small enough to sit in the handle's buffer until it closes, so a full disk surfaces + at close rather than at ``write()``, and the staging file left behind would still hold the token.""" + + def write_bytes_then_run_out_of_space(path: Path, data: bytes) -> int: + real_write_bytes(path, data) + raise OSError(errno.ENOSPC, "No space left on device") + + monkeypatch.setattr(Path, "write_bytes", write_bytes_then_run_out_of_space) + store = FileTokenStore(tmp_path) + + store.save( + "closing", + StoredToken(access_token=SecretStr("sk-ant-oat01-live"), expires_at_epoch=None, assertion_sha256="sha"), + ) + + leaked = [path for path in tmp_path.rglob("*") if path.is_file() and "sk-ant-oat01-live" in path.read_text()] + assert leaked == [], "a staged token file survived the close that failed" + assert store.load("closing") is None + + +def test_a_rotated_assertion_buys_a_fresh_token_that_other_workers_pick_up(tmp_path: Path): + poster = SingleUsePoster() + store = FileTokenStore(tmp_path) + assertions = {DEFAULT_REF: "jwt-v1"} + first = minted(store_engine(poster, store, reader=assertions).get_token(make_spec())) + + assertions[DEFAULT_REF] = "jwt-v2" + rotated = minted(store_engine(poster, store, reader=assertions).get_token(make_spec())) + follower = minted(store_engine(poster, store, reader=assertions).get_token(make_spec())) + + assert rotated.access_token.get_secret_value() != first.access_token.get_secret_value() + assert follower.access_token.get_secret_value() == rotated.access_token.get_secret_value() + assert [request["assertion"] for request in poster.requests] == ["jwt-v1", "jwt-v2"] + + +def test_an_expired_shared_token_is_not_reused(tmp_path: Path): + poster = ScriptedPoster([token_response("first", expires_in=60), token_response("second", expires_in=60)]) + store = FileTokenStore(tmp_path) + wall = FakeClock(1_700_000_000.0) + minted(store_engine(poster, store, wall_clock=wall).get_token(make_spec())) + + wall.advance(61) + later = minted(store_engine(poster, store, wall_clock=wall).get_token(make_spec())) + + assert later.access_token.get_secret_value() == "second" + assert len(poster.requests) == 2 + + +def test_the_remaining_lifetime_survives_different_monotonic_origins(tmp_path: Path): + poster = ScriptedPoster([token_response(expires_in=3600)]) + store = FileTokenStore(tmp_path) + wall = FakeClock(1_700_000_000.0) + minted(store_engine(poster, store, clock=FakeClock(1_000.0), wall_clock=wall).get_token(make_spec())) + + wall.advance(600) + later_clock = FakeClock(50_000.0) + later = minted(store_engine(poster, store, clock=later_clock, wall_clock=wall).get_token(make_spec())) + + assert later.expires_at == pytest.approx(50_000.0 + 3000.0) + assert len(poster.requests) == 1 + + +def test_mandatory_refresh_serves_the_shared_token_until_it_expires_then_fails_once(tmp_path: Path): + """With an unrotated assertion there is nothing new to exchange: refreshes inside the mandatory + window keep serving the shared token, and once it has expired the one allowed POST is denied + without a second identical POST behind it.""" + poster = SingleUsePoster(expires_in=3600) + store = FileTokenStore(tmp_path) + clock = FakeClock(1_000.0) + engine = store_engine(poster, store, clock=clock, wall_clock=clock) + first = minted(engine.get_token(make_spec())) + + clock.advance(3600 - 20) + refreshed = minted(engine.get_token(make_spec())) + assert refreshed.access_token.get_secret_value() == first.access_token.get_secret_value() + assert len(poster.requests) == 1 + + clock.advance(25) + failed = engine.get_token(make_spec()) + + assert isinstance(failed, TokenEndpointError) + assert failed.status_code == 401 + assert len(poster.requests) == 2 + + +def test_a_corrupt_cache_entry_is_treated_as_absent(tmp_path: Path): + poster = ScriptedPoster([token_response("first"), token_response("second")]) + store = FileTokenStore(tmp_path) + minted(store_engine(poster, store).get_token(make_spec())) + (entry,) = stored_files(tmp_path) + entry.write_text("{not json") + + later = minted(store_engine(poster, store).get_token(make_spec())) + + assert later.access_token.get_secret_value() == "second" + assert json.loads(entry.read_text())["access_token"] == "second" + + +def test_cache_entries_are_private_to_the_owner(tmp_path: Path): + store = FileTokenStore(tmp_path / "cache") + minted(store_engine(ScriptedPoster([token_response()]), store).get_token(make_spec())) + + (entry,) = stored_files(tmp_path / "cache") + assert stat.S_IMODE((tmp_path / "cache").stat().st_mode) == 0o700 + assert stat.S_IMODE(entry.stat().st_mode) == 0o600 + + +def test_a_group_readable_cache_directory_is_refused_and_the_engine_still_mints(tmp_path: Path): + loose = tmp_path / "loose" + loose.mkdir(mode=0o750) + os.chmod(loose, 0o750) + poster = ScriptedPoster([token_response("first"), token_response("second")]) + store = FileTokenStore(loose) + + minted(store_engine(poster, store).get_token(make_spec())) + later = minted(store_engine(poster, store).get_token(make_spec())) + + assert later.access_token.get_secret_value() == "second" + assert stored_files(loose) == [] + + +def test_default_store_follows_the_cache_dir_env(tmp_path: Path, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv(CACHE_DIR_ENV, "") + assert default_shared_token_store() is None + + monkeypatch.setenv(CACHE_DIR_ENV, str(tmp_path / "configured")) + configured = default_shared_token_store() + assert isinstance(configured, FileTokenStore) + assert configured.directory == tmp_path / "configured" + + monkeypatch.delenv(CACHE_DIR_ENV) + default = default_shared_token_store() + assert isinstance(default, FileTokenStore) + assert default.directory.name == f"litellm-token-exchange-{os.getuid()}" + + +class GatedSingleUsePoster(SingleUsePoster): + def __init__(self) -> None: + super().__init__() + self.entered = threading.Event() + self.release = threading.Event() + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.entered.set() + assert self.release.wait(timeout=10) + return super().post(url, content=content, headers=headers, timeout=timeout) + + +def test_a_worker_arriving_mid_exchange_waits_for_the_leader_instead_of_posting(tmp_path: Path): + poster = GatedSingleUsePoster() + store = FileTokenStore(tmp_path) + leader = store_engine(poster, store) + follower = store_engine(poster, store) + results: dict[str, object] = {} + + def lead() -> None: + results["leader"] = leader.get_token(make_spec()) + + def follow() -> None: + results["follower"] = follower.get_token(make_spec()) + + leader_thread = threading.Thread(target=lead, daemon=True) + leader_thread.start() + assert poster.entered.wait(timeout=10) + follower_thread = threading.Thread(target=follow, daemon=True) + follower_thread.start() + follower_thread.join(timeout=0.5) + assert follower_thread.is_alive() + poster.release.set() + leader_thread.join(timeout=10) + follower_thread.join(timeout=10) + + assert not follower_thread.is_alive() + assert ( + minted(results["follower"]).access_token.get_secret_value() + == minted(results["leader"]).access_token.get_secret_value() + ) + assert len(poster.requests) == 1 + + +def test_invalidate_drops_the_shared_entry(tmp_path: Path): + poster = ScriptedPoster([token_response("first"), token_response("second")]) + store = FileTokenStore(tmp_path) + engine = store_engine(poster, store) + spec: Final = make_spec() + minted(engine.get_token(spec)) + + engine.invalidate(spec) + + assert stored_files(tmp_path) == [] + assert minted(store_engine(poster, store).get_token(spec)).access_token.get_secret_value() == "second" diff --git a/tests/unit/llms/base_llm/auth/test_token_exchange.py b/tests/unit/llms/base_llm/auth/test_token_exchange.py new file mode 100644 index 00000000000..52b4837ba97 --- /dev/null +++ b/tests/unit/llms/base_llm/auth/test_token_exchange.py @@ -0,0 +1,1807 @@ +import asyncio +import concurrent.futures +import base64 +import json +from urllib.parse import quote, urlencode +import logging +import threading +import time +from collections.abc import Callable, Mapping +from types import MappingProxyType +from typing import Final +from urllib.parse import parse_qsl + +import httpx +import pytest +from pydantic import SecretStr + +from litellm.llms.base_llm.auth.token_exchange import ( + _METRICS_QUEUE_LIMIT, + _REDACTION_CAP, + ADVISORY_REFRESH_BACKOFF_SECONDS, + CALL_TYPE_CACHE_HIT, + FALLBACK_TOKEN_TTL_SECONDS, + MAX_ASSERTION_BYTES, + MAX_RESPONSE_BYTES, + JwtBearerTokenExchangeEngine, + ServiceLoggingMetricsSink, + TokenExchangeEndpointFailure, + TokenExchangeTransportFailure, + _default_assertion_reader, + _error_summary, + _HttpxSyncTokenPoster, + _new_exchange_handler, + redact_oauth_error_body, +) +from litellm.llms.base_llm.auth.types import ( + AssertionSource, + AssertionSourceError, + BodyEncoding, + ExchangeError, + ExchangeResult, + InsecureTokenUrl, + MalformedTokenResponse, + MintedToken, + TokenEndpointError, + TokenExchangeSpec, + TokenTransportError, +) +from litellm.secret_managers.main import OidcPathNotAllowedError, _resolve_oidc_file_path +from litellm.types.services import ServiceTypes + +DEFAULT_REF: Final = "oidc/env/TEST_ASSERTION" +DEFAULT_ASSERTION: Final = "test-jwt-assertion" +EXCHANGE_URL: Final = "https://token.example/v1/oauth/token" + + +class FakeClock: + def __init__(self, start: float = 1_000.0) -> None: + self.now = start + + def __call__(self) -> float: + return self.now + + def advance(self, seconds: float) -> None: + self.now += seconds + + +class RecordedRequest: + def __init__(self, url: str, content: bytes, headers: Mapping[str, str], timeout: float) -> None: + self.url = url + self.content = content + self.headers = dict(headers) + self.timeout = timeout + + def json_body(self) -> dict: + return json.loads(self.content) + + +class ScriptedPoster: + """Returns scripted responses in order (repeating the last one); records requests.""" + + def __init__( + self, + responses: list[httpx.Response], + on_request: Callable[[RecordedRequest], None] | None = None, + ) -> None: + self.requests: list[RecordedRequest] = [] + self._responses = list(responses) + self._on_request = on_request + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + recorded = RecordedRequest(url, content, headers, timeout) + self.requests.append(recorded) + if self._on_request is not None: + self._on_request(recorded) + if len(self._responses) > 1: + return self._responses.pop(0) + return self._responses[0] + + +class RaisingPoster: + def __init__(self, error: Exception) -> None: + self.calls = 0 + self._error = error + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.calls += 1 + raise self._error + + +class ManualExecutor(concurrent.futures.Executor): + """Records submissions; runs them only when the test says so.""" + + def __init__(self) -> None: + self.pending: list[Callable[[], None]] = [] + + def submit(self, fn, /, *args, **kwargs): + future: concurrent.futures.Future = concurrent.futures.Future() + self.pending.append(lambda: fn(*args, **kwargs)) + return future + + def run_all(self) -> None: + drained = list(self.pending) + self.pending.clear() + for job in drained: + job() + + +class InlineExecutor(concurrent.futures.Executor): + def submit(self, fn, /, *args, **kwargs): + future: concurrent.futures.Future = concurrent.futures.Future() + future.set_result(fn(*args, **kwargs)) + return future + + +class NeverRunsExecutor(concurrent.futures.Executor): + """Accepts work and never runs it, standing in for a telemetry backend that has stalled, so a + test can show the backlog stops growing instead of consuming memory for as long as traffic lasts.""" + + def __init__(self) -> None: + self.submitted = 0 # mutable-ok: a test spy counting accepted work + + def submit(self, fn, /, *args, **kwargs): + self.submitted += 1 + return concurrent.futures.Future() + + +class RefusingExecutor(concurrent.futures.Executor): + """A pool that has already been shut down, which is what an advisory refresh finds when the + worker is on its way out and a request still lands on a cached identity.""" + + def submit(self, fn, /, *args, **kwargs): + raise RuntimeError("cannot schedule new futures after shutdown") + + +class RecordingReader: + """Reports which identity-source refs were actually read off disk or out of the environment.""" + + def __init__(self) -> None: + self.reads: list[str] = [] + + def __call__(self, ref: str) -> str | None: + self.reads.append(ref) + return DEFAULT_ASSERTION + + +def token_response(token: str = "sk-ant-oat01-minted", expires_in: int | None = 3600) -> httpx.Response: + body: Final[dict[str, str | int]] = { + "access_token": token, + "token_type": "Bearer", + **({} if expires_in is None else {"expires_in": expires_in}), + } + return httpx.Response(200, json=body) + + +def make_spec( + *, + token_url: str = "https://token.example/v1/oauth/token", + assertion_ref: str = DEFAULT_REF, + assertion_field: str = "assertion", + static_body: Mapping[str, str] = MappingProxyType( + { + "grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer", + "federation_rule_id": "fdrl_1", + "organization_id": "org-1", + } + ), + body_encoding: BodyEncoding = "json", + request_headers: Mapping[str, str] = MappingProxyType( + {"anthropic-beta": "oauth-2025-04-20,oidc-federation-2026-04-01"} + ), + cache_key_identity: tuple[str, ...] = ("fdrl_1", "org-1", "", ""), + timeout_seconds: float = 2.0, + assertion_source: AssertionSource | None = None, +) -> TokenExchangeSpec: + return TokenExchangeSpec( + token_url=token_url, + assertion_ref=assertion_ref, + assertion_field=assertion_field, + static_body=static_body, + body_encoding=body_encoding, + request_headers=request_headers, + cache_key_identity=cache_key_identity, + timeout_seconds=timeout_seconds, + assertion_source=assertion_source, + ) + + +class RecordingMetricsSink: + def __init__(self) -> None: + self.successes: list[tuple[str, float]] = [] + self.failures: list[tuple[str, float, ExchangeError]] = [] + self.cache_hits = 0 + + def exchange_success(self, *, call_type: str, duration_seconds: float) -> None: + self.successes.append((call_type, duration_seconds)) + + def exchange_failure(self, *, call_type: str, duration_seconds: float, error: ExchangeError) -> None: + self.failures.append((call_type, duration_seconds, error)) + + def cache_hit(self) -> None: + self.cache_hits += 1 + + +def make_engine( + poster, + reader: Mapping[str, str] | Callable[[str], str | None] | None = None, + clock: FakeClock | None = None, + executor: concurrent.futures.Executor | None = None, + max_entries: int = 64, + metrics_sink=None, +) -> JwtBearerTokenExchangeEngine: + resolved_reader = reader if callable(reader) else (reader or {DEFAULT_REF: DEFAULT_ASSERTION}).get + return JwtBearerTokenExchangeEngine( + poster=poster, + assertion_reader=resolved_reader, + clock=clock if clock is not None else FakeClock(), + refresh_executor=executor if executor is not None else ManualExecutor(), + max_entries=max_entries, + metrics_sink=metrics_sink if metrics_sink is not None else RecordingMetricsSink(), + ) + + +def mint(engine: JwtBearerTokenExchangeEngine, spec: TokenExchangeSpec) -> MintedToken: + result = engine.get_token(spec) + assert isinstance(result, MintedToken) + return result + + +class TestFreshMintWireExact: + def test_json_body_and_headers(self): + poster = ScriptedPoster([token_response(expires_in=3600)]) + clock = FakeClock(start=1_000.0) + engine = make_engine(poster, clock=clock) + spec = make_spec() + + result = mint(engine, spec) + + assert result.access_token.get_secret_value() == "sk-ant-oat01-minted" + assert result.expires_at == 1_000.0 + 3600 + assert len(poster.requests) == 1 + request = poster.requests[0] + assert request.url == "https://token.example/v1/oauth/token" + assert request.timeout == 2.0 + assert request.headers == { + "content-type": "application/json", + "anthropic-beta": "oauth-2025-04-20,oidc-federation-2026-04-01", + } + assert request.json_body() == { + "grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer", + "federation_rule_id": "fdrl_1", + "organization_id": "org-1", + "assertion": DEFAULT_ASSERTION, + } + + def test_form_body_and_content_type(self): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + spec = make_spec(body_encoding="form") + + mint(engine, spec) + + request = poster.requests[0] + assert request.headers["content-type"] == "application/x-www-form-urlencoded" + assert dict(parse_qsl(request.content.decode())) == { + "grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer", + "federation_rule_id": "fdrl_1", + "organization_id": "org-1", + "assertion": DEFAULT_ASSERTION, + } + + +def test_cache_hit_zero_posts(): + poster = ScriptedPoster([token_response(expires_in=3600)]) + clock = FakeClock() + engine = make_engine(poster, clock=clock) + spec = make_spec() + + first = mint(engine, spec) + clock.advance(100.0) + second = mint(engine, spec) + + assert len(poster.requests) == 1 + assert second.access_token.get_secret_value() == first.access_token.get_secret_value() + + +@pytest.mark.parametrize( + "remaining,expect_advisory_submit,expect_new_token", + [ + (121.0, False, False), + (120.0, True, False), + (119.0, True, False), + (31.0, True, False), + (30.0, False, True), + (29.0, False, True), + ], +) +def test_window_boundaries(remaining: float, expect_advisory_submit: bool, expect_new_token: bool): + poster = ScriptedPoster([token_response("old-token", expires_in=3600), token_response("new-token")]) + clock = FakeClock(start=1_000.0) + executor = ManualExecutor() + engine = make_engine(poster, clock=clock, executor=executor) + spec = make_spec() + + mint(engine, spec) + expires_at = 1_000.0 + 3600 + clock.now = expires_at - remaining + result = mint(engine, spec) + + assert len(executor.pending) == (1 if expect_advisory_submit else 0) + expected_token = "new-token" if expect_new_token else "old-token" + assert result.access_token.get_secret_value() == expected_token + assert len(poster.requests) == (2 if expect_new_token else 1) + + +def test_advisory_serve_stale_single_flight_backoff(caplog: pytest.LogCaptureFixture): + poster = ScriptedPoster( + [ + token_response("stale-token", expires_in=3600), + httpx.Response(500, json={"error": "server_error"}), + httpx.Response(500, json={"error": "server_error"}), + ] + ) + clock = FakeClock(start=1_000.0) + executor = ManualExecutor() + engine = make_engine(poster, clock=clock, executor=executor) + spec = make_spec() + + mint(engine, spec) + clock.now = 1_000.0 + 3600 - 100.0 + + first = mint(engine, spec) + second = mint(engine, spec) + assert first.access_token.get_secret_value() == "stale-token" + assert second.access_token.get_secret_value() == "stale-token" + assert len(executor.pending) == 1 + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + executor.run_all() + assert len(poster.requests) == 2 + warning_records = [r for r in caplog.records if r.levelno == logging.WARNING] + assert any("Advisory token refresh" in r.getMessage() for r in warning_records) + assert "server_error" in caplog.text + assert DEFAULT_ASSERTION not in caplog.text + assert "stale-token" not in caplog.text + + within_backoff = mint(engine, spec) + assert within_backoff.access_token.get_secret_value() == "stale-token" + assert len(executor.pending) == 0 + + clock.advance(ADVISORY_REFRESH_BACKOFF_SECONDS) + after_backoff = mint(engine, spec) + assert after_backoff.access_token.get_secret_value() == "stale-token" + assert len(executor.pending) == 1 + executor.run_all() + assert len(poster.requests) == 3 + + +def test_an_advisory_refresh_the_executor_refuses_still_leaves_the_identity_mintable(): + """The advisory path arms the entry before handing the work off, so an executor that refuses the + submit used to strand it as in-flight with nothing on its way to publish: the cached token kept + serving until it expired, and every call after that waited out the follower timeout and failed, + so that identity never minted again.""" + poster = ScriptedPoster([token_response("first-token", expires_in=3600), token_response("second-token")]) + clock = FakeClock(start=1_000.0) + engine = make_engine(poster, clock=clock, executor=RefusingExecutor()) + spec = make_spec() + + mint(engine, spec) + clock.now = 1_000.0 + 3600 - 100.0 + inside_advisory_window = mint(engine, spec) + + assert inside_advisory_window.access_token.get_secret_value() == "first-token" + assert len(poster.requests) == 1 + + clock.now = 1_000.0 + 3600 + 1.0 + after_expiry = mint(engine, spec) + + assert after_expiry.access_token.get_secret_value() == "second-token" + assert len(poster.requests) == 2 + + +class GatedPoster: + """Blocks the leader inside post() until the test releases it.""" + + def __init__(self, response: httpx.Response) -> None: + self.entered = threading.Event() + self.release = threading.Event() + self.calls = 0 + self._calls_lock = threading.Lock() + self._response = response + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + with self._calls_lock: + self.calls += 1 + self.entered.set() + assert self.release.wait(timeout=10) + return self._response + + +def _run_concurrent_get_token( + engine: JwtBearerTokenExchangeEngine, spec: TokenExchangeSpec, poster: GatedPoster, thread_count: int +) -> list[ExchangeResult]: + results: list[ExchangeResult] = [] + results_lock = threading.Lock() + start_barrier = threading.Barrier(thread_count) + + def worker() -> None: + start_barrier.wait() + result = engine.get_token(spec) + with results_lock: + results.append(result) + + threads = [threading.Thread(target=worker, daemon=True) for _ in range(thread_count)] + for thread in threads: + thread.start() + assert poster.entered.wait(timeout=10) + time.sleep(0.3) + poster.release.set() + for thread in threads: + thread.join(timeout=10) + assert not thread.is_alive() + return results + + +def test_mandatory_single_leader(): + poster = GatedPoster(token_response("leader-token")) + engine = make_engine(poster) + spec = make_spec() + + results = _run_concurrent_get_token(engine, spec, poster, thread_count=5) + + assert poster.calls == 1 + assert len(results) == 5 + for result in results: + assert isinstance(result, MintedToken) + assert result.access_token.get_secret_value() == "leader-token" + + +def test_mandatory_failure_is_value(): + poster = GatedPoster(httpx.Response(500, json={"error": "server_error"})) + engine = make_engine(poster) + spec = make_spec() + + results = _run_concurrent_get_token(engine, spec, poster, thread_count=3) + + assert len(results) == 3 + for result in results: + assert isinstance(result, TokenEndpointError) + assert result.status_code == 500 + assert "server_error" in result.redacted_body + + +def test_lock_released_around_io(): + inner_spec = make_spec( + token_url="https://inner.example/v1/oauth/token", + cache_key_identity=("fdrl_inner", "org-1", "", ""), + ) + engine_holder: dict[str, JwtBearerTokenExchangeEngine] = {} + inner_results: list[ExchangeResult] = [] + + class ReentrantPoster: + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + if url == "https://token.example/v1/oauth/token": + inner_results.append(engine_holder["engine"].get_token(inner_spec)) + return token_response() + + engine = make_engine(ReentrantPoster()) + engine_holder["engine"] = engine + + outcome: list[ExchangeResult] = [] + thread = threading.Thread(target=lambda: outcome.append(engine.get_token(make_spec())), daemon=True) + thread.start() + thread.join(timeout=10) + + assert not thread.is_alive(), "engine held its lock across poster I/O and deadlocked" + assert len(outcome) == 1 + assert isinstance(outcome[0], MintedToken) + assert len(inner_results) == 1 + assert isinstance(inner_results[0], MintedToken) + + +def test_401_retry_once_with_reread(): + assertions = {DEFAULT_REF: "assertion-v1"} + + def rotate_on_first_request(request: RecordedRequest) -> None: + assertions[DEFAULT_REF] = "assertion-v2" + + poster = ScriptedPoster( + [httpx.Response(401, json={"error": "invalid_grant"}), token_response()], + on_request=rotate_on_first_request, + ) + engine = make_engine(poster, reader=assertions.get) + + result = mint(engine, make_spec()) + + assert result.access_token.get_secret_value() == "sk-ant-oat01-minted" + assert len(poster.requests) == 2 + assert poster.requests[0].json_body()["assertion"] == "assertion-v1" + assert poster.requests[1].json_body()["assertion"] == "assertion-v2" + + +class RotatingAssertionSource: + """A per-call assertion source that mints a fresh value on every read -- the shape + internal_issuer/keycloak identity sources take (a fresh JWT/token minted per call).""" + + def __init__(self, values: list[str]) -> None: + self._values = iter(values) + self.calls = 0 + + def __call__(self) -> str: + self.calls += 1 + return next(self._values) + + +class EchoingUnauthorizedPoster: + """401s every attempt, echoing the submitted assertion back into the error body -- a + token endpoint that reflects the request.""" + + def __init__(self) -> None: + self.requests: list[RecordedRequest] = [] + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + recorded = RecordedRequest(url, content, headers, timeout) + self.requests.append(recorded) + submitted = recorded.json_body()["assertion"] + return httpx.Response(401, json={"error": "invalid_grant", "error_description": f"bad assertion {submitted}"}) + + +def test_401_retry_redacts_the_assertion_actually_sent_not_a_fresh_reread(): + """Regression: with a rotating identity source, the reflection-drop check must match the + assertion the failing (second) attempt actually sent. Re-reading for the check would mint a + THIRD value that was never sent, so the reflection probe would miss and the actually-sent, + actually-reflected second assertion would leak into the error.""" + poster = EchoingUnauthorizedPoster() + source = RotatingAssertionSource(["assertion-v1", "assertion-v2", "assertion-v3"]) + engine = make_engine(poster) + spec = make_spec(assertion_source=source) + + result = engine.get_token(spec) + + assert isinstance(result, TokenEndpointError) + assert len(poster.requests) == 2 + assert poster.requests[0].json_body()["assertion"] == "assertion-v1" + assert poster.requests[1].json_body()["assertion"] == "assertion-v2" + assert source.calls == 2, "the failing attempt's own assertion must be reused, never re-read a third time" + assert "assertion-v1" not in result.redacted_body + assert "assertion-v2" not in result.redacted_body + assert "assertion-v3" not in result.redacted_body + + +def test_401_with_an_unchanged_assertion_is_not_resent(): + """An issuer that consumed the assertion's ``jti`` denies the identical assertion again, so the + retry only happens when the re-read assertion differs from the one the 401 came back for.""" + poster = ScriptedPoster([httpx.Response(401, json={"error": "invalid_grant"})]) + engine = make_engine(poster) + + result = engine.get_token(make_spec()) + + assert isinstance(result, TokenEndpointError) + assert result.status_code == 401 + assert "invalid_grant" in result.redacted_body + assert len(poster.requests) == 1 + + +class TestRedactionAndCaps: + def test_object_body_reduced_to_rfc6749_fields(self): + poster = ScriptedPoster( + [ + httpx.Response( + 400, + json={ + "error": "invalid_grant", + "error_description": "d" * 500, + "error_uri": "https://errors.example/e1", + "assertion_echo": "LEAKED-ASSERTION", + }, + ) + ] + ) + result = make_engine(poster).get_token(make_spec()) + + assert isinstance(result, TokenEndpointError) + assert result.status_code == 400 + assert "invalid_grant" in result.redacted_body + assert "d" * 256 in result.redacted_body + assert "d" * 257 not in result.redacted_body + assert "https://errors.example/e1" in result.redacted_body + assert "LEAKED-ASSERTION" not in result.redacted_body + + def test_nested_error_envelope_renders_readable_text(self): + body = { + "type": "error", + "error": { + "type": "invalid_request_error", + "message": "federation_rule_id is not a well-formed fdrl_ tagged ID", + }, + } + result = redact_oauth_error_body(400, json.dumps(body)) + + assert "invalid_request_error" in result.redacted_body + assert "federation_rule_id is not a well-formed fdrl_ tagged ID" in result.redacted_body + assert "{'" not in result.redacted_body + + def test_flat_rfc6749_shape_still_renders(self): + body = {"error": "invalid_grant", "error_description": "bad request"} + result = redact_oauth_error_body(400, json.dumps(body)) + + assert result.redacted_body == "error: invalid_grant; error_description: bad request" + + def test_nested_error_message_is_capped_at_256_chars(self): + body = {"error": {"type": "invalid_request_error", "message": "m" * 500}} + result = redact_oauth_error_body(400, json.dumps(body)) + + assert "m" * 256 in result.redacted_body + assert "m" * 257 not in result.redacted_body + + def test_json_string_body_is_not_echoed(self): + """A free-text body can carry back whatever was sent, so only structured OAuth fields are + ever rendered into an error an operator or caller will see.""" + result = redact_oauth_error_body(400, json.dumps("s" * 500)) + assert result.redacted_body == "non-object error response omitted" + assert "s" * 32 not in result.redacted_body + + def test_plain_text_body_is_not_echoed(self): + result = redact_oauth_error_body(502, "t" * 500) + assert result.redacted_body == "non-JSON error response omitted" + assert "t" * 32 not in result.redacted_body + + def test_reflected_assertion_is_dropped(self): + """An endpoint that echoes the submitted assertion must not put it in the log or the error.""" + assertion = SecretStr("eyJhbGciOiJSUzI1NiJ9.REFLECTEDPAYLOAD.signature") + body = {"error": "invalid_grant", "error_description": f"bad assertion {assertion.get_secret_value()}"} + + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert assertion.get_secret_value() not in result.redacted_body + assert "REFLECTEDPAYLOAD" not in result.redacted_body + + def test_assertion_reflected_from_an_offset_is_dropped(self): + """Regression: the probe only looked at the assertion's first 24 characters, so an + endpoint echoing it from any later offset shared no prefix and slipped through.""" + assertion = SecretStr("eyJhbGciOiJSUzI1NiJ9." + "A" * 40 + "PAYLOADMIDDLE" + "B" * 40 + ".signature") + tail = assertion.get_secret_value()[24:] + body = {"error": "invalid_grant", "error_description": tail} + + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert "PAYLOADMIDDLE" not in result.redacted_body + assert tail[:40] not in result.redacted_body + + def test_a_secret_carrying_spaces_is_dropped_when_echoed_whole(self): + """Regression on the redactor itself: comparing a compacted response against an + uncompacted secret stopped matching hand-set passphrases, which are exactly the secrets + most likely to be echoed and the ones an earlier contiguous match had caught.""" + assertion = SecretStr("correct horse battery staple, 42!") + echoed = assertion.get_secret_value() + body = {"error": "invalid_client", "error_description": f"secret {echoed} rejected"} + + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert echoed not in result.redacted_body + + def test_a_percent_encoded_secret_is_dropped(self): + """A form-encoded grant puts the secret on the wire percent-escaped, so an echo of that + shape has to be recognised without every caller enumerating it.""" + assertion = SecretStr("sUp3r+S3cret/Value=123") + echoed = quote(assertion.get_secret_value(), safe="") + body = {"error": "invalid_client", "error_description": f"rejected {echoed}"} + + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert echoed not in result.redacted_body + + def test_a_space_encoded_as_plus_is_dropped(self): + """A form-encoded body writes a space as "+", not %20, so percent-decoding alone does not + recover the secret and a passphrase echoed in its wire shape would travel on.""" + assertion = SecretStr("correct horse battery staple") + echoed = urlencode({"client_secret": assertion.get_secret_value()}).split("=", 1)[1] + body = {"error": "invalid_client", "error_description": f"rejected {echoed}"} + + assert "+" in echoed + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert echoed not in result.redacted_body + + def test_several_wire_forms_are_all_compared(self): + """The caller declares each shape it sent, since an encoding the redactor cannot reverse + (base64 of id:secret) is only knowable there.""" + raw = SecretStr("sUp3rS3cretValue123") + blob = SecretStr(base64.b64encode(b"litellm:sUp3rS3cretValue123").decode()) + body = {"error": "invalid_client", "error_description": f"bad {blob.get_secret_value()}"} + + result = redact_oauth_error_body(400, json.dumps(body), (raw, blob)) + + assert blob.get_secret_value() not in result.redacted_body + + def test_a_fragment_shorter_than_a_long_run_is_dropped(self): + """A slice too short to share a long contiguous run with the assertion is still assertion + material, and repeated errors would hand it over piece by piece.""" + assertion = SecretStr("eyJhbGciOiJSUzI1NiJ9." + "A" * 60 + ".sigsigsig") + fragment = assertion.get_secret_value()[30:48] + body = {"error": "invalid_grant", "error_description": f"rejected near {fragment}"} + + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert fragment not in result.redacted_body + + def test_a_fragment_broken_up_by_delimiters_is_dropped(self): + """Splitting the echo defeats a contiguous match, so the comparison ignores whatever the + endpoint put between the pieces.""" + assertion = SecretStr("eyJhbGciOiJSUzI1NiJ9." + "A" * 60 + ".sigsigsig") + piece = assertion.get_secret_value()[20:44] + spaced = " ".join(piece[i : i + 6] for i in range(0, 24, 6)) + body = {"error": "invalid_grant", "error_description": spaced} + + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert spaced not in result.redacted_body + + def test_a_short_secret_is_still_matched_whole(self): + """A Keycloak client secret can be shorter than the probe length; the whole value is + compared in that case rather than a truncated prefix.""" + secret = SecretStr("short-secret") + body = {"error": "invalid_client", "error_description": "rejected short-secret"} + + result = redact_oauth_error_body(400, json.dumps(body), secret) + + assert "short-secret" not in result.redacted_body + + def test_a_short_secret_echoed_in_its_wire_shape_is_dropped(self): + """Regression: the run scan only ever compared eight-character windows, so a secret + with fewer credential characters than that could never match once it came back + percent-encoded rather than verbatim, and the whole-value check needs the raw form.""" + secret = SecretStr("p@ss w0rd!") + echoed = quote(secret.get_secret_value(), safe="") + body = {"error": "invalid_client", "error_description": f"rejected {echoed}"} + + assert secret.get_secret_value() not in echoed + result = redact_oauth_error_body(400, json.dumps(body), secret) + + assert echoed not in result.redacted_body + + def test_an_unrelated_body_is_not_falsely_redacted(self): + """The scan must not fire on a body that merely shares short runs with the assertion.""" + assertion = SecretStr("eyJhbGciOiJSUzI1NiJ9." + "Z" * 60 + ".signature") + body = {"error": "invalid_grant", "error_description": "the federation rule was not found"} + + result = redact_oauth_error_body(400, json.dumps(body), assertion) + + assert "the federation rule was not found" in result.redacted_body + + def test_json_array_body_constant_message(self): + result = redact_oauth_error_body(400, json.dumps(["a", "b"])) + assert result.redacted_body == "non-object error response omitted" + + def test_oversized_body_never_parsed(self): + poster = ScriptedPoster([httpx.Response(400, content=b'{"error": "' + b"x" * MAX_RESPONSE_BYTES + b'"}')]) + result = make_engine(poster).get_token(make_spec()) + + assert isinstance(result, TokenEndpointError) + assert result.redacted_body == "oversized error response omitted" + + def test_oversized_success_body_is_malformed(self): + poster = ScriptedPoster( + [httpx.Response(200, content=b'{"access_token": "' + b"x" * MAX_RESPONSE_BYTES + b'"}')] + ) + result = make_engine(poster).get_token(make_spec()) + + assert not isinstance(result, MintedToken) + assert b"x" * 10 not in str(result).encode() + + +@pytest.mark.parametrize("access_token", ["", " "]) +def test_empty_access_token_is_malformed(access_token: str): + poster = ScriptedPoster( + [httpx.Response(200, json={"access_token": access_token, "token_type": "Bearer", "expires_in": 3600})] + ) + result = make_engine(poster).get_token(make_spec()) + + assert isinstance(result, MalformedTokenResponse) + assert "empty access_token" in result.detail + + +def test_sentinel_leak_audit(caplog: pytest.LogCaptureFixture): + jwt_sentinel = "JWT-SENTINEL-2c9f1e7ab4" + token_sentinel = "sk-ant-oat01-TOKEN-SENTINEL-90d4c3aa17" + ref = "oidc/env/SENTINEL_ASSERTION" + + with caplog.at_level(logging.DEBUG): + success_poster = ScriptedPoster([token_response(token_sentinel, expires_in=3600)]) + success_clock = FakeClock() + success_executor = ManualExecutor() + engine = make_engine(success_poster, reader={ref: jwt_sentinel}, clock=success_clock, executor=success_executor) + spec = make_spec(assertion_ref=ref) + minted = mint(engine, spec) + + endpoint_error = make_engine( + ScriptedPoster([httpx.Response(400, json={"error": "invalid_grant"})]), reader={ref: jwt_sentinel} + ).get_token(spec) + transport_error = make_engine(RaisingPoster(RuntimeError("boom")), reader={ref: jwt_sentinel}).get_token(spec) + malformed_error = make_engine( + ScriptedPoster([httpx.Response(200, json={"unexpected": "shape"})]), reader={ref: jwt_sentinel} + ).get_token(spec) + oversized_error = make_engine( + ScriptedPoster([token_response()]), reader={ref: jwt_sentinel + "x" * MAX_ASSERTION_BYTES} + ).get_token(spec) + insecure_error = make_engine(ScriptedPoster([token_response()]), reader={ref: jwt_sentinel}).get_token( + make_spec(assertion_ref=ref, token_url="http://token.example/v1/oauth/token") + ) + + success_poster._responses = [httpx.Response(500, json={"error": "server_error"})] + success_clock.now = success_clock.now + 3600 - 100.0 + stale = engine.get_token(spec) + success_executor.run_all() + + audited_values = [ + str(minted), + repr(minted), + str(minted.access_token), + repr(minted.access_token), + str(endpoint_error), + repr(endpoint_error), + str(transport_error), + repr(transport_error), + str(malformed_error), + repr(malformed_error), + str(oversized_error), + repr(oversized_error), + str(insecure_error), + repr(insecure_error), + str(stale), + repr(stale), + caplog.text, + ] + assert isinstance(oversized_error, AssertionSourceError) + assert oversized_error.kind == "oversized" + for value in audited_values: + assert jwt_sentinel not in value + assert token_sentinel not in value + + +class TestAssertionGuards: + @pytest.mark.parametrize( + "assertion_value,expected_kind", + [ + ("x" * (MAX_ASSERTION_BYTES + 1), "oversized"), + (" \n\t ", "empty"), + (None, "missing"), + ], + ) + def test_bad_assertion_values(self, assertion_value: str | None, expected_kind: str): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster, reader=lambda ref: assertion_value) + + result = engine.get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.kind == expected_kind + assert result.source_ref == DEFAULT_REF + assert len(poster.requests) == 0 + + @pytest.mark.parametrize( + "raised,expected_kind", + [ + (OidcPathNotAllowedError("path outside allowed credential directories"), "disallowed_path"), + (ValueError("Environment variable ANTHROPIC_IDENTITY_TOKEN not found"), "unreadable"), + (ImportError("needs PyJWT and cryptography: pip install 'litellm[proxy]'"), "unreadable"), + (OSError("permission denied"), "unreadable"), + ], + ) + def test_raising_reader(self, raised: Exception, expected_kind: str): + poster = ScriptedPoster([token_response()]) + + def reader(ref: str) -> str | None: + raise raised + + result = make_engine(poster, reader=reader).get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.kind == expected_kind + assert len(poster.requests) == 0 + + def test_value_error_message_is_captured_as_detail(self): + poster = ScriptedPoster([token_response()]) + + def reader(ref: str) -> str | None: + raise ValueError("Keycloak token endpoint returned invalid_client") + + result = make_engine(poster, reader=reader).get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.detail == "Keycloak token endpoint returned invalid_client" + + def test_import_error_message_is_captured_as_detail(self): + poster = ScriptedPoster([token_response()]) + + def reader(ref: str) -> str | None: + raise ImportError("the internal_issuer identity source needs PyJWT and cryptography: pip install 'litellm[proxy]'") + + result = make_engine(poster, reader=reader).get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.detail is not None + assert "litellm[proxy]" in result.detail + + @pytest.mark.parametrize( + "raised", + [OidcPathNotAllowedError("path outside allowed credential directories"), OSError("permission denied")], + ) + def test_non_value_error_never_populates_detail(self, raised: Exception): + """Only the ValueError branch carries operator-diagnosable text; every other reader failure + stays detail=None, matching today's file/env behavior byte-for-byte.""" + poster = ScriptedPoster([token_response()]) + + def reader(ref: str) -> str | None: + raise raised + + result = make_engine(poster, reader=reader).get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.detail is None + + def test_value_error_detail_is_capped(self): + poster = ScriptedPoster([token_response()]) + overlong_message = "x" * (_REDACTION_CAP + 100) + + def reader(ref: str) -> str | None: + raise ValueError(overlong_message) + + result = make_engine(poster, reader=reader).get_token(make_spec()) + + assert isinstance(result, AssertionSourceError) + assert result.detail == overlong_message[:_REDACTION_CAP] + + +class TestAssertionSourceOverridesEngineReader: + """``TokenExchangeSpec.assertion_source`` is the dispatch mechanism a per-config identity + source (internal_issuer, keycloak) plugs into the shared engine with -- it must win over the + engine-level reader, and failures must still be reported against ``assertion_ref``.""" + + def test_assertion_source_is_used_instead_of_the_reader(self): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster, reader=lambda ref: "from-engine-reader") + spec = make_spec(assertion_source=lambda: "from-assertion-source") + + result = mint(engine, spec) + + assert result.access_token.get_secret_value() == "sk-ant-oat01-minted" + assert poster.requests[0].json_body()["assertion"] == "from-assertion-source" + + def test_reader_is_never_called_when_assertion_source_is_set(self): + poster = ScriptedPoster([token_response()]) + calls: list[str] = [] + + def reader(ref: str) -> str | None: + calls.append(ref) + return "from-engine-reader" + + engine = make_engine(poster, reader=reader) + spec = make_spec(assertion_source=lambda: "from-assertion-source") + + mint(engine, spec) + + assert calls == [] + + def test_assertion_source_failure_is_reported_against_assertion_ref(self): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster, reader=lambda ref: "from-engine-reader") + + def raising_source() -> str | None: + raise ValueError("keycloak token endpoint returned invalid_client") + + spec = make_spec(assertion_source=raising_source, assertion_ref="oidc/keycloak/abc123") + + result = engine.get_token(spec) + + assert isinstance(result, AssertionSourceError) + assert result.source_ref == "oidc/keycloak/abc123" + assert result.detail == "keycloak token endpoint returned invalid_client" + assert len(poster.requests) == 0 + + def test_assertion_source_is_re_invoked_on_401_retry(self): + """The retry's second attempt must also prefer ``assertion_source`` for the assertion it + sends, not silently fall back to the engine reader.""" + values = iter(["assertion-v1", "assertion-v2"]) + poster = ScriptedPoster([httpx.Response(401, json={"error": "invalid_grant"}), token_response()]) + engine = make_engine(poster, reader=lambda ref: "from-engine-reader") + spec = make_spec(assertion_source=lambda: next(values)) + + result = mint(engine, spec) + + assert result.access_token.get_secret_value() == "sk-ant-oat01-minted" + assert poster.requests[0].json_body()["assertion"] == "assertion-v1" + assert poster.requests[1].json_body()["assertion"] == "assertion-v2" + + +class TestOidcFilePathAllowlistRaisesTypedError: + """The engine classifies assertion-source failures by exception type (see + TestAssertionGuards.test_raising_reader); that classification only works if the real + oidc/file allowlist actually raises OidcPathNotAllowedError rather than a bare ValueError.""" + + def test_out_of_allowlist_absolute_path(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", raising=False) + + with pytest.raises(OidcPathNotAllowedError): + _resolve_oidc_file_path("/etc/not-a-credential-dir/token") + + def test_relative_path(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.delenv("LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", raising=False) + + with pytest.raises(OidcPathNotAllowedError): + _resolve_oidc_file_path("relative/token/path") + + +class TestHttpsEnforcement: + def test_plain_http_rejected_host_only_zero_posts(self): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + + result = engine.get_token(make_spec(token_url="http://token.example/v1/oauth/token")) + + assert result == InsecureTokenUrl(host="token.example") + assert "/v1/oauth/token" not in str(result) + assert len(poster.requests) == 0 + + def test_plain_http_rejected_before_the_assertion_is_read(self): + reader = RecordingReader() + engine = make_engine(ScriptedPoster([token_response()]), reader=reader) + + result = engine.get_token(make_spec(token_url="http://token.example/v1/oauth/token")) + + assert isinstance(result, InsecureTokenUrl) + assert reader.reads == [] + + @pytest.mark.parametrize( + "url", + [ + "http://localhost:8080/v1/oauth/token", + "http://127.0.0.1/v1/oauth/token", + "http://[::1]/v1/oauth/token", + ], + ) + def test_localhost_http_allowed(self, url: str): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + + result = engine.get_token(make_spec(token_url=url)) + + assert isinstance(result, MintedToken) + assert len(poster.requests) == 1 + + +def test_cache_key_semantics(): + poster = ScriptedPoster([token_response()]) + assertions = {DEFAULT_REF: DEFAULT_ASSERTION, "oidc/env/OTHER": "other-assertion"} + engine = make_engine(poster, reader=assertions.get) + base_spec = make_spec() + + mint(engine, base_spec) + mint(engine, make_spec(cache_key_identity=("fdrl_1", "org-1", "svc-2", ""))) + mint(engine, make_spec(token_url="https://other.example/v1/oauth/token")) + mint(engine, make_spec(assertion_ref="oidc/env/OTHER")) + assert len(poster.requests) == 4 + + assertions[DEFAULT_REF] = "rotated-assertion" + cached = mint(engine, base_spec) + assert len(poster.requests) == 4 + assert cached.access_token.get_secret_value() == "sk-ant-oat01-minted" + + +def test_the_cache_returns_to_its_bound_after_an_all_in_flight_burst(): + """An entry a leader owns is never evictable, so a burst of distinct identities can push the map + past max_entries. It must come back down once those entries are idle, rather than holding the + high-water mark for the life of the process.""" + clock = FakeClock() + engine = make_engine(ScriptedPoster([token_response(expires_in=3600)]), clock=clock, max_entries=4) + + def spec_for(index: int) -> TokenExchangeSpec: + return make_spec(cache_key_identity=("fdrl_1", f"org-{index}", "", "")) + + for index in range(12): + mint(engine, spec_for(index)) + + assert len(engine._entries) <= 4, ( # noqa: SLF001 # the bound under test is internal state + f"the cap is enforced once entries are idle, saw {len(engine._entries)}" + ) + + +def test_bounded_eviction(): + clock = FakeClock() + + class PerCallPoster: + def __init__(self) -> None: + self.calls = 0 + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.calls += 1 + body = json.loads(content) + expires_in = 3600 + int(body["organization_id"].split("-")[1]) + return token_response(f"token-{body['organization_id']}", expires_in=expires_in) + + poster = PerCallPoster() + engine = make_engine(poster, clock=clock, max_entries=64) + + def spec_for(index: int) -> TokenExchangeSpec: + return make_spec( + static_body={ + "grant_type": "urn:ietf:params:oauth:grant-type:jwt-bearer", + "federation_rule_id": "fdrl_1", + "organization_id": f"org-{index}", + }, + cache_key_identity=("fdrl_1", f"org-{index}", "", ""), + ) + + for index in range(65): + mint(engine, spec_for(index)) + assert poster.calls == 65 + + mint(engine, spec_for(0)) + assert poster.calls == 66, "the earliest-expiring entry (index 0) should have been evicted" + + mint(engine, spec_for(2)) + assert poster.calls == 66, "a later-expiring entry should still be cached" + + mint(engine, spec_for(1)) + assert poster.calls == 67, "re-inserting index 0 should have evicted the next earliest-expiring entry" + + +@pytest.mark.parametrize("expires_in", [None, 0, -5]) +def test_missing_or_nonsense_expires_in_gets_fallback_ttl(expires_in: int | None): + poster = ScriptedPoster( + [token_response("short-lived", expires_in=expires_in), token_response("reminted", expires_in=3600)] + ) + clock = FakeClock(start=1_000.0) + engine = make_engine(poster, clock=clock) + spec = make_spec() + + first = mint(engine, spec) + assert first.expires_at == 1_000.0 + FALLBACK_TOKEN_TTL_SECONDS + + clock.advance(FALLBACK_TOKEN_TTL_SECONDS + 1.0) + second = mint(engine, spec) + + assert second.access_token.get_secret_value() == "reminted" + assert len(poster.requests) == 2, "a token without a sane expires_in must never be cached forever" + + +async def test_aget_token_loop_responsive(): + class SleepingPoster: + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + time.sleep(0.3) + return token_response() + + engine = make_engine(SleepingPoster()) + spec = make_spec() + ticks = {"count": 0} + stop = asyncio.Event() + + async def ticker() -> None: + while not stop.is_set(): + ticks["count"] += 1 + await asyncio.sleep(0.01) + + ticker_task = asyncio.create_task(ticker()) + result = await engine.aget_token(spec) + stop.set() + await ticker_task + + assert isinstance(result, MintedToken) + assert result.access_token.get_secret_value() == "sk-ant-oat01-minted" + assert ticks["count"] >= 5, "the event loop was blocked during aget_token" + sync_result = engine.get_token(spec) + assert sync_result == result + + +def test_invalidate_forces_refresh(): + poster = ScriptedPoster([token_response("token-1", expires_in=3600), token_response("token-2", expires_in=3600)]) + engine = make_engine(poster) + spec = make_spec() + + first = mint(engine, spec) + assert first.access_token.get_secret_value() == "token-1" + + engine.invalidate(spec) + second = mint(engine, spec) + assert second.access_token.get_secret_value() == "token-2" + assert len(poster.requests) == 2 + + third = mint(engine, spec) + assert third.access_token.get_secret_value() == "token-2" + assert len(poster.requests) == 2, "force_refresh must be one-shot" + + +def test_invalidate_unknown_spec_is_noop(): + poster = ScriptedPoster([token_response()]) + engine = make_engine(poster) + + engine.invalidate(make_spec()) + + assert len(poster.requests) == 0 + + +def test_advisory_failure_wakes_expired_follower_to_re_lead(): + poster = ScriptedPoster( + [ + token_response("initial-token", expires_in=3600), + httpx.Response(500, json={"error": "server_error"}), + token_response("recovered-token", expires_in=3600), + ] + ) + clock = FakeClock(start=1_000.0) + executor = ManualExecutor() + engine = make_engine(poster, clock=clock, executor=executor) + spec = make_spec() + + mint(engine, spec) + clock.now = 1_000.0 + 3600 - 100.0 + mint(engine, spec) + assert len(executor.pending) == 1 + + clock.advance(200.0) + results: list[ExchangeResult] = [] + follower = threading.Thread(target=lambda: results.append(engine.get_token(spec)), daemon=True) + follower.start() + time.sleep(0.3) + executor.run_all() + follower.join(timeout=10) + + assert not follower.is_alive() + assert len(results) == 1 + result = results[0] + assert isinstance(result, MintedToken), f"follower was handed {result!r} instead of re-leading a fresh mint" + assert result.access_token.get_secret_value() == "recovered-token" + assert len(poster.requests) == 3 + + +class TwoAttemptGatedPoster: + """401 on the first attempt, then blocks the leader's retry until released.""" + + def __init__(self) -> None: + self.entered_second = threading.Event() + self.release = threading.Event() + self.calls = 0 + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.calls += 1 + if self.calls == 1: + return httpx.Response(401, json={"error": "invalid_grant"}) + self.entered_second.set() + assert self.release.wait(timeout=30) + return token_response("slow-leader-token") + + +def test_follower_budget_outlasts_slow_two_attempt_leader(): + poster = TwoAttemptGatedPoster() + rotating_reads = iter(["jwt-before-rotation", "jwt-after-rotation"]) + engine = make_engine(poster, reader=lambda ref: next(rotating_reads, "jwt-after-rotation")) + spec = make_spec(timeout_seconds=1.0) + + leader_results: list[ExchangeResult] = [] + leader = threading.Thread(target=lambda: leader_results.append(engine.get_token(spec)), daemon=True) + leader.start() + assert poster.entered_second.wait(timeout=10) + + follower_results: list[ExchangeResult] = [] + follower = threading.Thread(target=lambda: follower_results.append(engine.get_token(spec)), daemon=True) + follower.start() + time.sleep(6.5) + poster.release.set() + leader.join(timeout=10) + follower.join(timeout=10) + + assert leader_results and isinstance(leader_results[0], MintedToken) + assert follower_results, "follower never returned" + follower_result = follower_results[0] + assert isinstance(follower_result, MintedToken), ( + f"follower gave up before the leader's two-attempt worst case: {follower_result!r}" + ) + assert follower_result.access_token.get_secret_value() == "slow-leader-token" + + +class FailThenGatePoster: + """500 on the first call, then blocks until released before succeeding.""" + + def __init__(self) -> None: + self.entered_gate = threading.Event() + self.release = threading.Event() + self.calls = 0 + + def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: + self.calls += 1 + if self.calls == 1: + return httpx.Response(500, json={"error": "server_error"}) + self.entered_gate.set() + assert self.release.wait(timeout=30) + return token_response("round-two-token") + + +def test_new_round_timed_out_follower_never_returns_previous_rounds_error(): + poster = FailThenGatePoster() + clock = FakeClock() + engine = make_engine(poster, clock=clock) + spec = make_spec(timeout_seconds=0.05) + + first = engine.get_token(spec) + assert isinstance(first, TokenEndpointError) + + clock.advance(ADVISORY_REFRESH_BACKOFF_SECONDS + 1.0) + leader = threading.Thread(target=lambda: engine.get_token(spec), daemon=True) + leader.start() + assert poster.entered_gate.wait(timeout=10) + + follower_result = engine.get_token(spec) + + assert isinstance(follower_result, TokenTransportError), ( + f"timed-out follower returned the previous round's error: {follower_result!r}" + ) + assert "timed out" in follower_result.detail + poster.release.set() + leader.join(timeout=10) + + +def test_lead_backoff_fails_fast_within_window_and_expires_after(): + poster = ScriptedPoster([httpx.Response(500, json={"error": "server_error"})]) + clock = FakeClock() + engine = make_engine(poster, clock=clock) + spec = make_spec() + + first = engine.get_token(spec) + assert isinstance(first, TokenEndpointError) + assert len(poster.requests) == 1 + + clock.advance(ADVISORY_REFRESH_BACKOFF_SECONDS - 1.0) + second = engine.get_token(spec) + assert second == first + assert len(poster.requests) == 1, "a request inside the backoff window must make zero POSTs" + + clock.advance(1.0) + third = engine.get_token(spec) + assert isinstance(third, TokenEndpointError) + assert len(poster.requests) == 2 + + +def test_invalidate_bypasses_lead_backoff(): + poster = ScriptedPoster( + [httpx.Response(500, json={"error": "server_error"}), token_response("post-invalidate", expires_in=3600)] + ) + clock = FakeClock() + engine = make_engine(poster, clock=clock) + spec = make_spec() + + first = engine.get_token(spec) + assert isinstance(first, TokenEndpointError) + + engine.invalidate(spec) + second = engine.get_token(spec) + + assert isinstance(second, MintedToken) + assert second.access_token.get_secret_value() == "post-invalidate" + assert len(poster.requests) == 2 + + +class StubExchangeHandler: + """Stands in for the HTTPHandler the default poster builds, so the poster's own contract is + testable without a socket.""" + + def __init__(self, result: httpx.Response | Exception | None) -> None: + self.calls = 0 + self._result = result + + def post(self, url: str, *, content: bytes, headers: dict[str, str], timeout: float) -> httpx.Response | None: + self.calls += 1 + if isinstance(self._result, Exception): + raise self._result + return self._result + + +class TestDefaultTokenPoster: + def test_builds_its_handler_once_and_reuses_it(self): + built: list[StubExchangeHandler] = [] + + def factory() -> StubExchangeHandler: + handler = StubExchangeHandler(httpx.Response(200, json={"access_token": "t"})) + built.append(handler) + return handler + + poster: Final = _HttpxSyncTokenPoster(handler_factory=factory) # pyright: ignore[reportArgumentType] # StubExchangeHandler stands in for the legacy-untyped HTTPHandler + for _ in range(3): + poster.post(EXCHANGE_URL, content=b"", headers={}, timeout=1.0) + + assert len(built) == 1 + assert built[0].calls == 3 + + def test_the_real_handler_refuses_to_follow_redirects(self): + assert _new_exchange_handler().client.follow_redirects is False, ( + "a redirected exchange POST would replay the workload assertion to the redirect target" + ) + + def test_an_http_status_error_becomes_its_response(self): + response: Final = httpx.Response( + 401, json={"error": "invalid_grant"}, request=httpx.Request("POST", EXCHANGE_URL) + ) + poster: Final = _HttpxSyncTokenPoster( + handler_factory=lambda: StubExchangeHandler( # pyright: ignore[reportArgumentType] # StubExchangeHandler stands in for the legacy-untyped HTTPHandler + httpx.HTTPStatusError("boom", request=response.request, response=response) + ) + ) + + assert poster.post(EXCHANGE_URL, content=b"", headers={}, timeout=1.0).status_code == 401 + + def test_a_missing_response_is_a_transport_error(self): + poster: Final = _HttpxSyncTokenPoster(handler_factory=lambda: StubExchangeHandler(None)) # pyright: ignore[reportArgumentType] # StubExchangeHandler stands in for the legacy-untyped HTTPHandler + + with pytest.raises(httpx.TransportError): + poster.post(EXCHANGE_URL, content=b"", headers={}, timeout=1.0) + + +class TestDefaultAssertionReader: + def test_reads_through_litellm_secret_resolution(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setenv("WIF_ASSERTION_FOR_DEFAULT_READER", "header.payload.signature") + + assert _default_assertion_reader("os.environ/WIF_ASSERTION_FOR_DEFAULT_READER") == "header.payload.signature" + + def test_an_unset_reference_reads_as_none(self): + assert _default_assertion_reader("os.environ/DEFINITELY_NOT_SET_WIF_ASSERTION_REF") is None + + +class TestErrorSummary: + def test_every_error_variant_summarises_without_carrying_a_secret(self): + summaries: Final = { + _error_summary(AssertionSourceError(kind="unreadable", source_ref="oidc/file/x")), + _error_summary(InsecureTokenUrl(host="token.internal")), + _error_summary(TokenEndpointError(status_code=401, redacted_body="invalid_grant")), + _error_summary(TokenTransportError(detail="ConnectError: refused")), + _error_summary(MalformedTokenResponse(detail="empty access_token")), + } + + assert {s.split(":")[0] for s in summaries} == { + "AssertionSourceError", + "InsecureTokenUrl", + "TokenEndpointError", + "TokenTransportError", + "MalformedTokenResponse", + }, "each variant names itself so a log line says which stage failed" + + +class TestNonBearerTokenType: + def test_a_non_bearer_token_type_is_refused(self): + poster: Final = ScriptedPoster( + [httpx.Response(200, json={"access_token": "tok", "token_type": "mac", "expires_in": 300})] + ) + engine: Final = JwtBearerTokenExchangeEngine(poster=poster, assertion_reader=lambda _ref: DEFAULT_ASSERTION) + + result: Final = engine.get_token(make_spec()) + + assert isinstance(result, MalformedTokenResponse) + assert "non-bearer" in result.detail + + +class TestShortLivedRefreshWindows: + """A token whose lifetime is at or below the flat 120s advisory window used to be inside that + window from birth, so every request armed another background exchange. The windows now scale + with the observed lifetime; long-lived tokens must keep the flat 120s/30s behaviour.""" + + @staticmethod + def _engine_with( + expires_in: int | None, + ) -> tuple[JwtBearerTokenExchangeEngine, ScriptedPoster, FakeClock, ManualExecutor, TokenExchangeSpec]: + poster = ScriptedPoster([token_response("short-lived", expires_in=expires_in), token_response("reminted")]) + clock = FakeClock(start=1_000.0) + executor = ManualExecutor() + engine = make_engine(poster, clock=clock, executor=executor) + return engine, poster, clock, executor, make_spec() + + def test_fallback_ttl_token_is_served_without_arming_a_refresh(self): + engine, poster, clock, executor, spec = self._engine_with(expires_in=None) + + first = mint(engine, spec) + assert first.expires_at == 1_000.0 + FALLBACK_TOKEN_TTL_SECONDS + + for _ in range(5): + clock.advance(1.0) + assert mint(engine, spec).access_token.get_secret_value() == "short-lived" + + assert executor.pending == [], "a freshly minted fallback-TTL token must not arm a refresh on every request" + assert len(poster.requests) == 1 + + @pytest.mark.parametrize( + "elapsed,expect_advisory_submit", + [(29.0, False), (30.0, True), (52.0, True)], + ) + def test_fallback_ttl_token_refreshes_around_its_half_life(self, elapsed: float, expect_advisory_submit: bool): + engine, poster, clock, executor, spec = self._engine_with(expires_in=None) + + mint(engine, spec) + clock.advance(elapsed) + served = mint(engine, spec) + + assert served.access_token.get_secret_value() == "short-lived" + assert len(executor.pending) == (1 if expect_advisory_submit else 0) + executor.run_all() + assert len(poster.requests) == (2 if expect_advisory_submit else 1) + + @pytest.mark.parametrize( + "elapsed,expect_new_token", + [(52.0, False), (53.0, True)], + ) + def test_fallback_ttl_mandatory_wall_scales_with_the_lifetime(self, elapsed: float, expect_new_token: bool): + engine, poster, clock, executor, spec = self._engine_with(expires_in=None) + + mint(engine, spec) + clock.advance(elapsed) + served = mint(engine, spec) + + assert served.access_token.get_secret_value() == ("reminted" if expect_new_token else "short-lived") + assert len(executor.pending) == (0 if expect_new_token else 1) + + @pytest.mark.parametrize( + "elapsed,expect_advisory_submit", + [(89.0, False), (100.0, True)], + ) + def test_a_200s_token_scales_its_advisory_window_too(self, elapsed: float, expect_advisory_submit: bool): + engine, poster, clock, executor, spec = self._engine_with(expires_in=200) + + mint(engine, spec) + clock.advance(elapsed) + served = mint(engine, spec) + + assert served.access_token.get_secret_value() == "short-lived" + assert len(executor.pending) == (1 if expect_advisory_submit else 0) + + @pytest.mark.parametrize("expires_in", [240, 3600]) + @pytest.mark.parametrize( + "remaining,expect_advisory_submit,expect_new_token", + [ + (121.0, False, False), + (120.0, True, False), + (31.0, True, False), + (30.0, False, True), + ], + ) + def test_long_lived_tokens_keep_the_flat_windows( + self, expires_in: int, remaining: float, expect_advisory_submit: bool, expect_new_token: bool + ): + engine, poster, clock, executor, spec = self._engine_with(expires_in=expires_in) + + mint(engine, spec) + clock.now = 1_000.0 + expires_in - remaining + served = mint(engine, spec) + + assert len(executor.pending) == (1 if expect_advisory_submit else 0) + assert served.access_token.get_secret_value() == ("reminted" if expect_new_token else "short-lived") + assert len(poster.requests) == (2 if expect_new_token else 1) + + +class RaisingMetricsSink: + def exchange_success(self, *, call_type: str, duration_seconds: float) -> None: + raise RuntimeError("metrics sink down") + + def exchange_failure(self, *, call_type: str, duration_seconds: float, error: ExchangeError) -> None: + raise RuntimeError("metrics sink down") + + def cache_hit(self) -> None: + raise RuntimeError("metrics sink down") + + +class TestMetricsEmission: + def test_cold_mint_emits_success_with_duration(self): + clock = FakeClock() + sink = RecordingMetricsSink() + poster = ScriptedPoster([token_response()], on_request=lambda _request: clock.advance(0.25)) + engine = make_engine(poster, clock=clock, metrics_sink=sink) + + mint(engine, make_spec()) + + assert sink.successes == [("cold_mint", 0.25)] + assert sink.failures == [] + assert sink.cache_hits == 0 + + def test_cache_hit_emits_counter_not_a_mint(self): + clock = FakeClock() + sink = RecordingMetricsSink() + engine = make_engine(ScriptedPoster([token_response()]), clock=clock, metrics_sink=sink) + spec = make_spec() + + mint(engine, spec) + clock.advance(100.0) + mint(engine, spec) + + assert sink.cache_hits == 1 + assert len(sink.successes) == 1 + + def test_advisory_refresh_call_type(self): + clock = FakeClock(start=1_000.0) + sink = RecordingMetricsSink() + executor = ManualExecutor() + poster = ScriptedPoster([token_response("old", expires_in=3600), token_response("new")]) + engine = make_engine(poster, clock=clock, executor=executor, metrics_sink=sink) + spec = make_spec() + + mint(engine, spec) + clock.now = 1_000.0 + 3600 - 119.0 + mint(engine, spec) + executor.run_all() + + assert [call_type for call_type, _ in sink.successes] == ["cold_mint", "advisory_refresh"] + assert sink.cache_hits == 1 + + def test_mandatory_refresh_call_type(self): + clock = FakeClock(start=1_000.0) + sink = RecordingMetricsSink() + poster = ScriptedPoster([token_response("old", expires_in=3600), token_response("new")]) + engine = make_engine(poster, clock=clock, metrics_sink=sink) + spec = make_spec() + + mint(engine, spec) + clock.now = 1_000.0 + 3600 - 29.0 + mint(engine, spec) + + assert [call_type for call_type, _ in sink.successes] == ["cold_mint", "mandatory_refresh"] + assert sink.cache_hits == 0 + + def test_failed_exchange_emits_failure_once_and_negative_cache_does_not_reemit(self): + sink = RecordingMetricsSink() + poster = ScriptedPoster([httpx.Response(503, json={"error": "unavailable"})]) + engine = make_engine(poster, metrics_sink=sink) + spec = make_spec() + + first = engine.get_token(spec) + second = engine.get_token(spec) + + assert isinstance(first, TokenEndpointError) + assert isinstance(second, TokenEndpointError) + assert len(sink.failures) == 1 + call_type, _duration, error = sink.failures[0] + assert call_type == "cold_mint" + assert isinstance(error, TokenEndpointError) + assert error.status_code == 503 + assert sink.successes == [] + + def test_failure_payload_carries_no_assertion_material(self): + sink = RecordingMetricsSink() + engine = make_engine(EchoingUnauthorizedPoster(), metrics_sink=sink) + + result = engine.get_token(make_spec()) + + assert isinstance(result, TokenEndpointError) + (failure,) = sink.failures + assert DEFAULT_ASSERTION not in repr(failure) + assert DEFAULT_ASSERTION not in _error_summary(failure[2]) + + def test_raising_sink_never_breaks_mint_serve_or_failure(self): + clock = FakeClock() + engine = make_engine(ScriptedPoster([token_response()]), clock=clock, metrics_sink=RaisingMetricsSink()) + spec = make_spec() + + minted = mint(engine, spec) + clock.advance(100.0) + served = mint(engine, spec) + + assert served.access_token.get_secret_value() == minted.access_token.get_secret_value() + + failing = make_engine(RaisingPoster(httpx.ConnectError("boom")), metrics_sink=RaisingMetricsSink()) + result = failing.get_token(make_spec()) + assert isinstance(result, TokenTransportError) + + +class RecordingServiceHooks: + def __init__(self) -> None: + self.successes: list[tuple[ServiceTypes, str, float]] = [] + self.failures: list[tuple[ServiceTypes, float, str | Exception, str]] = [] + + async def async_service_success_hook(self, service: ServiceTypes, call_type: str, duration: float) -> None: + self.successes.append((service, call_type, duration)) + + async def async_service_failure_hook( + self, service: ServiceTypes, duration: float, error: str | Exception, call_type: str + ) -> None: + self.failures.append((service, duration, error, call_type)) + + +class RaisingServiceHooks: + """Every hook raises, and each call is recorded first so a test can prove the sink kept + calling through rather than bailing after the first failure.""" + + def __init__(self) -> None: + self.attempts: list[str] = [] # mutable-ok: a test spy accumulating calls in order + + async def async_service_success_hook(self, service: ServiceTypes, call_type: str, duration: float) -> None: + self.attempts.append(f"success:{call_type}") + raise RuntimeError("hook down") + + async def async_service_failure_hook( + self, service: ServiceTypes, duration: float, error: str | Exception, call_type: str + ) -> None: + self.attempts.append(f"failure:{call_type}") + raise RuntimeError("hook down") + + +class TestServiceLoggingMetricsSink: + def _sink(self, hooks) -> ServiceLoggingMetricsSink: + return ServiceLoggingMetricsSink(service_logging_factory=lambda: hooks, executor=InlineExecutor()) + + def test_success_maps_to_anthropic_wif_service(self): + hooks = RecordingServiceHooks() + + self._sink(hooks).exchange_success(call_type="cold_mint", duration_seconds=0.2) + + assert hooks.successes == [(ServiceTypes.ANTHROPIC_WIF, "cold_mint", 0.2)] + + def test_a_stalled_backend_stops_accepting_work_instead_of_queueing_without_bound(self): + stalled: Final = NeverRunsExecutor() + sink: Final = ServiceLoggingMetricsSink(service_logging_factory=RecordingServiceHooks, executor=stalled) + + for _ in range(_METRICS_QUEUE_LIMIT + 500): + sink.cache_hit() + + assert stalled.submitted == _METRICS_QUEUE_LIMIT, ( + "once the backlog is full further events are dropped, so request volume cannot grow it" + ) + + def test_a_drained_backlog_accepts_work_again(self): + hooks: Final = RecordingServiceHooks() + sink: Final = ServiceLoggingMetricsSink(service_logging_factory=lambda: hooks, executor=InlineExecutor()) + + for _ in range(_METRICS_QUEUE_LIMIT + 10): + sink.cache_hit() + + assert len(hooks.successes) == _METRICS_QUEUE_LIMIT + 10, ( + "an executor that actually runs releases each slot, so nothing is dropped" + ) + + def test_failure_maps_variant_and_redacted_summary(self): + hooks = RecordingServiceHooks() + error = TokenEndpointError(status_code=503, redacted_body="error: unavailable") + + self._sink(hooks).exchange_failure(call_type="mandatory_refresh", duration_seconds=0.1, error=error) + + ((service, duration, emitted, call_type),) = hooks.failures + assert service is ServiceTypes.ANTHROPIC_WIF + assert duration == 0.1 + assert call_type == "mandatory_refresh" + assert isinstance(emitted, TokenExchangeEndpointFailure) + assert str(emitted) == _error_summary(error) + + def test_transport_failure_gets_its_own_error_class(self): + hooks = RecordingServiceHooks() + + self._sink(hooks).exchange_failure( + call_type="advisory_refresh", duration_seconds=0.05, error=TokenTransportError(detail="ConnectError: boom") + ) + + ((_service, _duration, emitted, _call_type),) = hooks.failures + assert isinstance(emitted, TokenExchangeTransportFailure) + + def test_cache_hit_maps_to_cache_service_with_zero_duration(self): + hooks = RecordingServiceHooks() + + self._sink(hooks).cache_hit() + + assert hooks.successes == [(ServiceTypes.ANTHROPIC_WIF_CACHE, CALL_TYPE_CACHE_HIT, 0.0)] + + def test_end_to_end_reflected_assertion_never_reaches_the_hook(self): + hooks = RecordingServiceHooks() + sink = self._sink(hooks) + engine = make_engine(EchoingUnauthorizedPoster(), metrics_sink=sink) + + result = engine.get_token(make_spec()) + + assert isinstance(result, TokenEndpointError) + ((_service, _duration, emitted, call_type),) = hooks.failures + assert call_type == "cold_mint" + assert DEFAULT_ASSERTION not in str(emitted) + assert DEFAULT_ASSERTION not in repr(emitted) + + def test_raising_hooks_are_swallowed(self): + hooks: Final = RaisingServiceHooks() + sink: Final = self._sink(hooks) + + sink.exchange_success(call_type="cold_mint", duration_seconds=0.2) + sink.cache_hit() + sink.exchange_failure(call_type="cold_mint", duration_seconds=0.1, error=TokenTransportError(detail="boom")) + + assert hooks.attempts == ["success:cold_mint", "success:cache_hit", "failure:cold_mint"], ( + "every event is still handed to the hooks, and one raising hook does not stop the next" + ) diff --git a/tests/unit/llms/custom_httpx/test_http_handler.py b/tests/unit/llms/custom_httpx/test_http_handler.py index 15c842ade3e..af8cd6cbf24 100644 --- a/tests/unit/llms/custom_httpx/test_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_http_handler.py @@ -1247,6 +1247,53 @@ def test_sync_client_never_replays_one_upstreams_cookie_to_another(): assert seen == [None, None] +def _redirecting_upstream(): + """A host that answers every request with a redirect somewhere else, and records who was asked.""" + hosts = [] + + def handler(request: httpx.Request) -> httpx.Response: + hosts.append(request.url.host) + if request.url.host == "token.example": + return httpx.Response(302, headers={"location": "https://elsewhere.example/v1/oauth/token"}) + return httpx.Response(200, json={"access_token": "sk-ant-oat01-leaked"}) + + return handler, hosts + + +def test_a_handler_that_refuses_redirects_still_refuses_them_after_its_client_is_healed(): + """The token exchange handler refuses redirects because following one replays a signed identity + assertion at whatever host the Location header names. A closed client is healed by building a + fresh one, so a rebuild that read the setting off the code default rather than off the handler + would quietly start chasing them again for the rest of the process's life.""" + transport, hosts = _redirecting_upstream() + handler = HTTPHandler(follow_redirects=False) + handler.client._transport = httpx.MockTransport(transport) + + first = handler.client.get("https://token.example/v1/oauth/token") + handler.client.close() + + healed = handler.client + healed._transport = httpx.MockTransport(transport) + second = healed.get("https://token.example/v1/oauth/token") + + assert healed.is_closed is False + assert first.status_code == 302 + assert second.status_code == 302 + assert hosts == ["token.example", "token.example"] + + +def test_a_handler_left_on_the_default_still_follows_redirects(): + """Every other caller of the pool is an LLM provider call that has always followed redirects.""" + transport, hosts = _redirecting_upstream() + handler = HTTPHandler() + handler.client._transport = httpx.MockTransport(transport) + + response = handler.client.get("https://token.example/v1/oauth/token") + + assert response.status_code == 200 + assert hosts == ["token.example", "elsewhere.example"] + + @pytest.mark.asyncio async def test_aiohttp_session_never_replays_one_upstreams_cookie_to_another(): """The httpx jar is not the only one. AiohttpTransport is litellm's default transport diff --git a/tests/unit/llms/custom_httpx/test_llm_http_handler.py b/tests/unit/llms/custom_httpx/test_llm_http_handler.py index d283cc6c64c..21c62500df3 100644 --- a/tests/unit/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/unit/llms/custom_httpx/test_llm_http_handler.py @@ -22,7 +22,9 @@ from litellm.llms.base_llm.audio_transcription.transformation import ( AudioTranscriptionRequestData, BaseAudioTranscriptionConfig, ) +from litellm.llms.base_llm.batches.transformation import BaseBatchesConfig from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException +from litellm.llms.base_llm.files.transformation import BaseFilesConfig from litellm.llms.base_llm.search.transformation import BaseSearchConfig, SearchResponse from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS from litellm.llms.brave.search.transformation import BraveSearchConfig @@ -599,7 +601,7 @@ async def test_async_anthropic_messages_handler_extra_headers(): # Mock the config mock_config = Mock() - mock_config.validate_anthropic_messages_environment = Mock( + mock_config.avalidate_anthropic_messages_environment = AsyncMock( return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com") ) mock_config.transform_anthropic_messages_request = Mock( @@ -646,7 +648,7 @@ async def test_async_anthropic_messages_handler_extra_headers(): captured_headers.update(kwargs.get("headers", {})) return ({"x-api-key": "test-key"}, "https://api.anthropic.com") - mock_config.validate_anthropic_messages_environment = capture_validate + mock_config.avalidate_anthropic_messages_environment = AsyncMock(side_effect=capture_validate) try: await handler.async_anthropic_messages_handler( @@ -961,7 +963,7 @@ async def test_async_anthropic_messages_handler_passes_litellm_metadata(): handler = BaseLLMHTTPHandler() mock_config = Mock() - mock_config.validate_anthropic_messages_environment = Mock( + mock_config.avalidate_anthropic_messages_environment = AsyncMock( return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com") ) mock_config.transform_anthropic_messages_request = Mock( @@ -1040,7 +1042,7 @@ async def test_async_anthropic_messages_handler_forwards_router_model_info(): handler = BaseLLMHTTPHandler() mock_config = Mock() - mock_config.validate_anthropic_messages_environment = Mock( + mock_config.avalidate_anthropic_messages_environment = AsyncMock( return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com") ) mock_config.transform_anthropic_messages_request = Mock( @@ -1132,7 +1134,7 @@ async def test_async_anthropic_messages_handler_header_priority(): captured_headers.update(kwargs.get("headers", {})) return ({"x-api-key": "test-key"}, "https://api.anthropic.com") - mock_config.validate_anthropic_messages_environment = capture_validate + mock_config.avalidate_anthropic_messages_environment = AsyncMock(side_effect=capture_validate) mock_config.transform_anthropic_messages_request = Mock( return_value={"model": "claude-3-opus-20240229", "messages": []} ) @@ -1171,7 +1173,7 @@ async def test_async_anthropic_messages_handler_drops_top_level_and_nested_param handler = BaseLLMHTTPHandler() mock_config = Mock() - mock_config.validate_anthropic_messages_environment = Mock( + mock_config.avalidate_anthropic_messages_environment = AsyncMock( return_value=({"x-api-key": "test-key"}, "https://api.anthropic.com") ) @@ -1384,9 +1386,7 @@ def test_sync_delete_responses_sets_json_content_type(): ({}, True, None, None), ], ) -def test_resolve_anthropic_messages_timeout( - monkeypatch, litellm_params_kwargs, stream, global_timeout, expected -): +def test_resolve_anthropic_messages_timeout(monkeypatch, litellm_params_kwargs, stream, global_timeout, expected): from litellm.constants import DEFAULT_REQUEST_TIMEOUT_SECONDS if global_timeout is None: @@ -1402,9 +1402,7 @@ def test_resolve_anthropic_messages_timeout( ) else: monkeypatch.setattr("litellm.request_timeout", global_timeout, raising=False) - monkeypatch.setattr( - "litellm.request_timeout_explicitly_set", True, raising=False - ) + monkeypatch.setattr("litellm.request_timeout_explicitly_set", True, raising=False) resolved = BaseLLMHTTPHandler._resolve_anthropic_messages_timeout( litellm_params=GenericLiteLLMParams(**litellm_params_kwargs), @@ -1425,13 +1423,11 @@ async def test_async_anthropic_messages_handler_forwards_request_timeout(monkeyp handler = BaseLLMHTTPHandler() mock_config = Mock() - mock_config.validate_anthropic_messages_environment = Mock( + mock_config.avalidate_anthropic_messages_environment = AsyncMock( return_value=({"x-api-key": "k"}, "https://api.anthropic.com") ) mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False) - mock_config.transform_anthropic_messages_request = Mock( - return_value={"model": "claude", "messages": []} - ) + mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []}) mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages") mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None)) mock_config.max_retry_on_anthropic_messages_http_error = 1 @@ -1473,13 +1469,11 @@ async def test_async_anthropic_messages_handler_forwards_stream_timeout(monkeypa handler = BaseLLMHTTPHandler() mock_config = Mock() - mock_config.validate_anthropic_messages_environment = Mock( + mock_config.avalidate_anthropic_messages_environment = AsyncMock( return_value=({"x-api-key": "k"}, "https://api.anthropic.com") ) mock_config.should_filter_anthropic_beta_headers = Mock(return_value=False) - mock_config.transform_anthropic_messages_request = Mock( - return_value={"model": "claude", "messages": []} - ) + mock_config.transform_anthropic_messages_request = Mock(return_value={"model": "claude", "messages": []}) mock_config.get_complete_url = Mock(return_value="https://api.anthropic.com/v1/messages") mock_config.sign_request = Mock(return_value=({"x-api-key": "k"}, None)) mock_config.max_retry_on_anthropic_messages_http_error = 1 @@ -1881,7 +1875,7 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks( handler = BaseLLMHTTPHandler() mock_config = Mock() - mock_config.validate_anthropic_messages_environment = Mock( + mock_config.avalidate_anthropic_messages_environment = AsyncMock( return_value=({"x-api-key": "sk-test"}, "https://api.anthropic.com") ) mock_config.transform_anthropic_messages_request = Mock( @@ -1889,7 +1883,13 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks( ) mock_config.sign_request = Mock(return_value=({}, None)) - fake_raw_response = {"id": "msg_1", "type": "message", "role": "assistant", "content": [], "stop_reason": "end_turn"} + fake_raw_response = { + "id": "msg_1", + "type": "message", + "role": "assistant", + "content": [], + "stop_reason": "end_turn", + } mock_config.transform_anthropic_messages_response = Mock(return_value=fake_raw_response) mock_logging_obj = Mock() @@ -1909,10 +1909,17 @@ async def test_async_anthropic_messages_handler_passes_api_key_to_agentic_hooks( mock_httpx_response.status_code = 200 with ( - patch.object(handler, "_async_post_anthropic_messages_with_http_error_retry", new=AsyncMock(return_value=mock_httpx_response)), + patch.object( + handler, + "_async_post_anthropic_messages_with_http_error_retry", + new=AsyncMock(return_value=mock_httpx_response), + ), patch.object(handler, "_call_agentic_completion_hooks", side_effect=fake_agentic_hooks), patch("litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client"), - patch("litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers", return_value=None), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.litellm_core_utils.get_provider_specific_headers.ProviderSpecificHeaderUtils.get_provider_specific_headers", + return_value=None, + ), ): result = await handler.async_anthropic_messages_handler( model="claude-haiku", @@ -2245,7 +2252,9 @@ def test_audio_transcriptions_sends_dict_data_as_json_body(): form-encodes it and silently ignores json=; JSON-body providers (e.g. Google Speech-to-Text) need an application/json body.""" captured = {} - client = HTTPHandler(client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured)))) + client = HTTPHandler( + client=httpx.Client(transport=httpx.MockTransport(_capture_json_transcription_request(captured))) + ) response = BaseLLMHTTPHandler().audio_transcriptions( client=client, @@ -2404,6 +2413,105 @@ def test_sync_retrieve_file_content_raises_on_http_error(): assert exc_info.value.status_code == 404 +_FILE_CONTENT_WIF_ENV = { + "ANTHROPIC_FEDERATION_RULE_ID": "fdrl_llm_http_handler_seam", + "ANTHROPIC_ORGANIZATION_ID": "org-llm-http-handler-seam", + "ANTHROPIC_IDENTITY_TOKEN": "llm-http-handler-seam-inline-jwt", +} + + +class _BlockingWifPoster: + """A token-endpoint poster that blocks until released, so the test can prove + the exchange ran off the event loop's own thread instead of freezing it.""" + + def __init__(self): + self.release = threading.Event() + self.thread_ids = [] + + def post(self, url, *, content, headers, timeout): + self.thread_ids.append(threading.get_ident()) + self.release.wait(timeout=5) + return httpx.Response( + 200, + json={ + "access_token": "sk-ant-oat01-llm-http-handler-seam", + "token_type": "Bearer", + "expires_in": 3600, + }, + ) + + +@pytest.mark.asyncio +async def test_async_retrieve_file_content_wif_exchange_does_not_block_event_loop(monkeypatch): + """Regression (Greptile P1): async_retrieve_file_content called the synchronous + validate_environment directly, so a cold WIF mint on this call site froze the + event loop until the exchange finished. It must resolve credentials through the + async facade instead.""" + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.files.transformation import AnthropicFilesConfig + from litellm.llms.anthropic.wif import aget_anthropic_wif_token, get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + for name in ("ANTHROPIC_API_KEY", "ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_BASE", "ANTHROPIC_BASE_URL"): + monkeypatch.delenv(name, raising=False) + for name, value in _FILE_CONTENT_WIF_ENV.items(): + monkeypatch.setenv(name, value) + + poster = _BlockingWifPoster() + engine = JwtBearerTokenExchangeEngine(poster=poster) + sync_calls = [] + + def sync_shim(litellm_params, api_base, model): + sync_calls.append(model) + return get_anthropic_wif_token(litellm_params, api_base, model, engine) + + async def async_shim(litellm_params, api_base, model): + return await aget_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "get_anthropic_wif_token", sync_shim) + monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim) + + handler = BaseLLMHTTPHandler() + client = Mock(spec=AsyncHTTPHandler) + client.get = AsyncMock(return_value=httpx.Response(status_code=200, content=b"file bytes")) + + ticks = [] + + async def ticker(): + for i in range(20): + await asyncio.sleep(0.005) + ticks.append(i) + + ticker_task = asyncio.create_task(ticker()) + await asyncio.sleep(0.02) + + retrieve_task = asyncio.create_task( + handler.async_retrieve_file_content( + file_content_request={"file_id": "file-abc"}, + provider_config=AnthropicFilesConfig(), + litellm_params={}, + headers={}, + logging_obj=Mock(), + client=client, + ) + ) + await asyncio.sleep(0.05) + # The ticker kept advancing while the token exchange was still blocked on + # poster.release, proving the exchange did not run inline on the event loop. + assert len(ticks) > 0 + assert not retrieve_task.done() + + poster.release.set() + await retrieve_task + await ticker_task + + assert sync_calls == [] + assert poster.thread_ids + assert poster.thread_ids[0] != threading.get_ident() + sent_headers = client.get.call_args.kwargs["headers"] + assert sent_headers["authorization"] == "Bearer sk-ant-oat01-llm-http-handler-seam" + + _UPSTREAM_NOT_FOUND_BODY = { "error": { "message": "Response with id 'resp_abc' not found.", @@ -2551,9 +2659,7 @@ async def test_anthropic_invalid_thinking_signature_retry_resigns_bedrock_reques ok_response = httpx.Response(200, json={"id": "msg_1"}, request=httpx.Request("POST", request_url)) class FakeAsyncClient: - async def post( - self, url, headers, data, stream=False, logging_obj=None, timeout=None - ): + async def post(self, url, headers, data, stream=False, logging_obj=None, timeout=None): posts.append({"headers": dict(headers), "data": data}) return invalid_signature_response if len(posts) == 1 else ok_response @@ -2932,7 +3038,7 @@ async def test_async_anthropic_messages_handler_carries_deployment_vertex_locati custom_llm_provider="vertex_ai", ) mock_config = Mock() - mock_config.validate_anthropic_messages_environment = Mock( + mock_config.avalidate_anthropic_messages_environment = AsyncMock( return_value=({"authorization": "Bearer t"}, "https://us-east5-aiplatform.googleapis.com") ) mock_config.transform_anthropic_messages_request = Mock( @@ -3448,6 +3554,123 @@ async def test_a_provider_that_keeps_rejecting_is_not_retried_forever_on_the_asy assert len(recorder.bodies) == 2 +def _async_client_returning(response: Mock) -> AsyncMock: + client = AsyncMock(spec=AsyncHTTPHandler) + client.post.return_value = response + return client + + +@pytest.mark.asyncio +async def test_create_file_async_awaits_the_provider_credential_hook_instead_of_blocking(): + provider_config = Mock(spec=BaseFilesConfig) + provider_config.validate_environment.side_effect = AssertionError("sync validate_environment ran on the event loop") + provider_config.avalidate_environment = AsyncMock(return_value={"x-api-key": "federated"}) + provider_config.get_complete_file_url.return_value = "https://files.example/v1/files" + provider_config.transform_create_file_request.return_value = {"file": ("batch.jsonl", b"{}", "application/jsonl")} + file_object = object() + provider_config.transform_create_file_response.return_value = file_object + client = _async_client_returning(Mock(spec=httpx.Response)) + + result = await BaseLLMHTTPHandler().create_file( + create_file_data={"file": b"{}", "purpose": "batch"}, + litellm_params={}, + provider_config=provider_config, + headers={}, + api_base=None, + api_key=None, + logging_obj=Mock(), + _is_async=True, + client=client, + ) + + assert result is file_object + provider_config.validate_environment.assert_not_called() + provider_config.avalidate_environment.assert_awaited_once() + assert client.post.call_args.kwargs["headers"] == {"x-api-key": "federated"} + assert client.post.call_args.kwargs["url"] == "https://files.example/v1/files" + + +@pytest.mark.asyncio +async def test_create_batch_async_validates_credentials_off_the_event_loop(): + provider_config = Mock(spec=BaseBatchesConfig) + provider_config.validate_environment.side_effect = lambda **_: {"x-validated-on": str(threading.get_ident())} + provider_config.get_complete_batch_url.return_value = "https://batches.example/v1/messages/batches" + provider_config.transform_create_batch_request.return_value = {"requests": []} + batch = object() + provider_config.transform_create_batch_response.return_value = batch + client = _async_client_returning(Mock(spec=httpx.Response)) + + result = await BaseLLMHTTPHandler().create_batch( + create_batch_data={"input_file_id": "file_1", "endpoint": "/v1/chat/completions", "completion_window": "24h"}, + litellm_params={}, + provider_config=provider_config, + headers={}, + api_base=None, + api_key=None, + logging_obj=Mock(), + _is_async=True, + client=client, + model="claude-sonnet-4-5", + ) + + assert result is batch + validated_on = client.post.call_args.kwargs["headers"]["x-validated-on"] + assert validated_on != str(threading.get_ident()) + assert client.post.call_args.kwargs["url"] == "https://batches.example/v1/messages/batches" + + +@pytest.mark.asyncio +async def test_create_file_says_which_setting_is_missing_when_the_provider_resolves_no_url(): + """A provider that cannot work out where its files endpoint lives returns no URL, which used to + be posted as-is: the caller saw an httpx error about an invalid URL and no mention of api_base.""" + provider_config = Mock(spec=BaseFilesConfig) + provider_config.avalidate_environment = AsyncMock(return_value={}) + provider_config.get_complete_file_url.return_value = None + client = _async_client_returning(Mock(spec=httpx.Response)) + + with pytest.raises(ValueError, match="api_base is required for create_file"): + await BaseLLMHTTPHandler().create_file( + create_file_data={"file": b"{}", "purpose": "batch"}, + litellm_params={}, + provider_config=provider_config, + headers={}, + api_base=None, + api_key=None, + logging_obj=Mock(), + _is_async=True, + client=client, + ) + + client.post.assert_not_called() + provider_config.transform_create_file_request.assert_not_called() + + +@pytest.mark.asyncio +async def test_create_batch_says_which_setting_is_missing_when_the_provider_resolves_no_url(): + """Same on the batches path, which resolves its URL the same way.""" + provider_config = Mock(spec=BaseBatchesConfig) + provider_config.validate_environment.return_value = {} + provider_config.get_complete_batch_url.return_value = None + client = _async_client_returning(Mock(spec=httpx.Response)) + + with pytest.raises(ValueError, match="api_base is required for create_batch"): + await BaseLLMHTTPHandler().create_batch( + create_batch_data={"input_file_id": "file_1", "endpoint": "/v1/chat/completions", "completion_window": "24h"}, + litellm_params={}, + provider_config=provider_config, + headers={}, + api_base=None, + api_key=None, + logging_obj=Mock(), + _is_async=True, + client=client, + model="claude-sonnet-4-5", + ) + + client.post.assert_not_called() + provider_config.transform_create_batch_request.assert_not_called() + + CONTAINER_NOT_FOUND_BODY = { "error": { "message": "Container with id 'cntr_gone' not found.", diff --git a/tests/unit/llms/openai/test_openai_workload_identity.py b/tests/unit/llms/openai/test_openai_workload_identity.py index db107e00df0..74415d45638 100644 --- a/tests/unit/llms/openai/test_openai_workload_identity.py +++ b/tests/unit/llms/openai/test_openai_workload_identity.py @@ -10,6 +10,7 @@ from openai import AsyncOpenAI, OpenAI import litellm from litellm.llms.litellm_proxy.responses.transformation import LiteLLMProxyResponsesAPIConfig +from litellm.llms.openai.chat.gpt_transformation import OpenAIGPTConfig from litellm.llms.openai.common_utils import BaseOpenAILLM, OpenAIError from litellm.llms.openai.openai import OpenAIChatCompletion from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig @@ -22,6 +23,17 @@ from litellm.llms.openai.workload_identity import ( from litellm.types.router import GenericLiteLLMParams TOKEN_EXCHANGE_URL: Final = "https://auth.openai.com/oauth/token" +CHAT_COMPLETIONS_URL: Final = "https://api.openai.com/v1/chat/completions" +EMBEDDINGS_URL: Final = "https://api.openai.com/v1/embeddings" +MODELS_URL: Final = "https://api.openai.com/v1/models" +CHAT_COMPLETION_BODY: Final = { + "id": "chatcmpl-wif", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, +} @pytest.fixture @@ -279,3 +291,342 @@ class TestResponsesValidateEnvironment: headers={}, model="gpt-4o-mini", litellm_params=GenericLiteLLMParams() ) assert headers["Authorization"] == "Bearer None" + + +@pytest.fixture +def deployment_wif(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> dict[str, str]: + token_file: Final = tmp_path / "deployment_subject_token.jwt" + token_file.write_text("subject-token-from-deployment-file") + for name in ( + "OPENAI_API_KEY", + "OPENAI_BASE_URL", + "OPENAI_API_BASE", + "OPENAI_IDENTITY_PROVIDER_ID", + "OPENAI_SERVICE_ACCOUNT_ID", + "OPENAI_IDENTITY_TOKEN_FILE", + ): + monkeypatch.delenv(name, raising=False) + monkeypatch.setattr(litellm, "api_base", None) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + _workload_identity_auth.cache_clear() + litellm.in_memory_llm_clients_cache.flush_cache() + return { + "openai_identity_provider_id": "idp_deployment", + "openai_service_account_id": "user-deployment", + "openai_identity_token_file": str(token_file), + } + + +def deployment_config(deployment_wif: dict[str, str]) -> OpenAIWorkloadIdentityConfig: + return OpenAIWorkloadIdentityConfig( + identity_provider_id="idp_deployment", + service_account_id="user-deployment", + token_file=deployment_wif["openai_identity_token_file"], + ) + + +def mock_chat_completions() -> respx.Route: + return respx.post(CHAT_COMPLETIONS_URL).mock(return_value=httpx.Response(200, json=CHAT_COMPLETION_BODY)) + + +def mock_streaming_chat_completions() -> respx.Route: + chunk: Final = {"id": "chatcmpl-1", "object": "chat.completion.chunk", "created": 1, "model": "gpt-4o-mini"} + events: Final = ( + {**chunk, "choices": [{"index": 0, "delta": {"role": "assistant", "content": "ok"}, "finish_reason": None}]}, + {**chunk, "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + ) + body: Final = "".join(f"data: {json.dumps(event)}\n\n" for event in events) + "data: [DONE]\n\n" + return respx.post(CHAT_COMPLETIONS_URL).mock( + return_value=httpx.Response(200, headers={"content-type": "text/event-stream"}, content=body) + ) + + +class TestResolveConfigFromDeployment: + def test_resolves_from_litellm_params_without_env(self, deployment_wif: dict[str, str]) -> None: + assert resolve_openai_workload_identity_config( + api_key=None, api_base=None, litellm_params=deployment_wif + ) == deployment_config(deployment_wif) + + def test_env_alone_disables_nothing_when_params_are_absent(self, deployment_wif: dict[str, str]) -> None: + assert resolve_openai_workload_identity_config(api_key=None, api_base=None, litellm_params=None) is None + + def test_unrelated_litellm_params_do_not_resolve(self, deployment_wif: dict[str, str]) -> None: + assert ( + resolve_openai_workload_identity_config(api_key=None, api_base=None, litellm_params={"model": "gpt-4o"}) + is None + ) + + def test_litellm_params_beat_env(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: + config: Final = resolve_openai_workload_identity_config( + api_key=None, + api_base=None, + litellm_params={ + "openai_identity_provider_id": "idp_deployment", + "openai_service_account_id": "user-deployment", + "openai_identity_token_file": wif_env.token_file, + }, + ) + assert config == OpenAIWorkloadIdentityConfig( + identity_provider_id="idp_deployment", + service_account_id="user-deployment", + token_file=wif_env.token_file, + ) + + def test_partial_litellm_params_fill_from_env_per_field(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: + config: Final = resolve_openai_workload_identity_config( + api_key=None, api_base=None, litellm_params={"openai_identity_provider_id": "idp_deployment"} + ) + assert config == OpenAIWorkloadIdentityConfig( + identity_provider_id="idp_deployment", + service_account_id=wif_env.service_account_id, + token_file=wif_env.token_file, + ) + + @pytest.mark.parametrize("blank", ["", None, 7]) + def test_blank_or_non_string_param_falls_back_to_env( + self, wif_env: OpenAIWorkloadIdentityConfig, blank: object + ) -> None: + config: Final = resolve_openai_workload_identity_config( + api_key=None, api_base=None, litellm_params={"openai_identity_provider_id": blank} + ) + assert config == wif_env + + def test_partial_litellm_params_without_env_disable(self, deployment_wif: dict[str, str]) -> None: + partial: Final = {key: value for key, value in deployment_wif.items() if key != "openai_identity_token_file"} + assert resolve_openai_workload_identity_config(api_key=None, api_base=None, litellm_params=partial) is None + + def test_static_api_key_beats_litellm_params(self, deployment_wif: dict[str, str]) -> None: + assert ( + resolve_openai_workload_identity_config(api_key="sk-static", api_base=None, litellm_params=deployment_wif) + is None + ) + + def test_env_openai_api_key_beats_litellm_params( + self, deployment_wif: dict[str, str], monkeypatch: pytest.MonkeyPatch + ) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "sk-from-env") + assert ( + resolve_openai_workload_identity_config(api_key=None, api_base=None, litellm_params=deployment_wif) is None + ) + + def test_foreign_api_base_disables_deployment_wif(self, deployment_wif: dict[str, str]) -> None: + assert ( + resolve_openai_workload_identity_config( + api_key=None, api_base="https://my-vllm.internal/v1", litellm_params=deployment_wif + ) + is None + ) + + +class TestDeploymentClientConstruction: + def test_sync_client_from_deployment_params(self, deployment_wif: dict[str, str]) -> None: + client: Final = OpenAIChatCompletion()._get_openai_client( + is_async=False, api_key=None, api_base=None, litellm_params=deployment_wif + ) + assert isinstance(client, OpenAI) + assert client.api_key == "workload-identity-auth" + assert client._workload_identity_auth is not None + + def test_async_client_from_deployment_params(self, deployment_wif: dict[str, str]) -> None: + client: Final = OpenAIChatCompletion()._get_openai_client( + is_async=True, api_key=None, api_base=None, litellm_params=deployment_wif + ) + assert isinstance(client, AsyncOpenAI) + assert client._workload_identity_auth is not None + + def test_distinct_deployments_get_distinct_cached_clients(self, deployment_wif: dict[str, str]) -> None: + other_deployment: Final = {**deployment_wif, "openai_service_account_id": "user-other"} + handler: Final = OpenAIChatCompletion() + first: Final = handler._get_openai_client( + is_async=False, api_key=None, api_base=None, litellm_params=deployment_wif + ) + second: Final = handler._get_openai_client( + is_async=False, api_key=None, api_base=None, litellm_params=other_deployment + ) + again: Final = handler._get_openai_client( + is_async=False, api_key=None, api_base=None, litellm_params=dict(deployment_wif) + ) + assert first is not second + assert again is first + + @respx.mock + def test_completion_kwargs_carry_exchanged_bearer(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("deployment-bearer") + completion_route: Final = mock_chat_completions() + + response: Final = litellm.completion( + model="openai/gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], **deployment_wif + ) + + assert response.choices[0].message.content == "ok" + request: Final = completion_route.calls.last.request + assert request.headers["Authorization"] == "Bearer deployment-bearer" + assert not any(key.startswith("openai_") for key in json.loads(request.content)) + + @respx.mock + def test_streaming_completion_kwargs_carry_exchanged_bearer(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("stream-bearer") + stream_route: Final = mock_streaming_chat_completions() + + chunks: Final = tuple( + litellm.completion( + model="openai/gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], stream=True, **deployment_wif + ) + ) + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert stream_route.calls.last.request.headers["Authorization"] == "Bearer stream-bearer" + + @respx.mock + @pytest.mark.asyncio + async def test_async_streaming_completion_kwargs_carry_exchanged_bearer( + self, deployment_wif: dict[str, str] + ) -> None: + mock_token_exchange("async-stream-bearer") + stream_route: Final = mock_streaming_chat_completions() + + stream: Final = await litellm.acompletion( + model="openai/gpt-4o-mini", messages=[{"role": "user", "content": "hi"}], stream=True, **deployment_wif + ) + chunks: Final = tuple([chunk async for chunk in stream]) + + assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "ok" + assert stream_route.calls.last.request.headers["Authorization"] == "Bearer async-stream-bearer" + + @respx.mock + def test_router_deployment_without_api_key_authenticates_via_token_exchange( + self, deployment_wif: dict[str, str] + ) -> None: + exchange_route: Final = mock_token_exchange("router-bearer") + completion_route: Final = mock_chat_completions() + router: Final = litellm.Router( + model_list=[{"model_name": "wif-gpt", "litellm_params": {"model": "openai/gpt-4o-mini", **deployment_wif}}] + ) + + response: Final = router.completion(model="wif-gpt", messages=[{"role": "user", "content": "hi"}]) + + assert response.choices[0].message.content == "ok" + assert exchange_route.called + assert completion_route.calls.last.request.headers["Authorization"] == "Bearer router-bearer" + + @respx.mock + def test_embedding_kwargs_carry_exchanged_bearer(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("embedding-bearer") + embeddings_route: Final = respx.post(EMBEDDINGS_URL).mock( + return_value=httpx.Response( + 200, + json={ + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2]}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 1, "total_tokens": 1}, + }, + ) + ) + + litellm.embedding(model="openai/text-embedding-3-small", input=["hi"], **deployment_wif) + + assert embeddings_route.calls.last.request.headers["Authorization"] == "Bearer embedding-bearer" + + +class TestResponsesValidateEnvironmentFromDeployment: + @respx.mock + def test_mints_bearer_from_litellm_params(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("responses-bearer") + headers: Final = OpenAIResponsesAPIConfig().validate_environment( + headers={}, model="gpt-4o-mini", litellm_params=GenericLiteLLMParams(**deployment_wif) + ) + assert headers["Authorization"] == "Bearer responses-bearer" + + def test_static_key_in_litellm_params_wins(self, deployment_wif: dict[str, str]) -> None: + headers: Final = OpenAIResponsesAPIConfig().validate_environment( + headers={}, + model="gpt-4o-mini", + litellm_params=GenericLiteLLMParams(api_key="sk-responses", **deployment_wif), + ) + assert headers["Authorization"] == "Bearer sk-responses" + + +class TestDiscoverModels: + @staticmethod + def mock_models() -> respx.Route: + return respx.get(MODELS_URL).mock( + return_value=httpx.Response(200, json={"data": [{"id": "gpt-4o-mini"}, {"id": "gpt-4.1"}]}) + ) + + @respx.mock + def test_discovers_with_exchanged_bearer_from_litellm_params(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("discovery-bearer") + models_route: Final = self.mock_models() + + assert OpenAIGPTConfig().discover_models(deployment_wif) == ["gpt-4o-mini", "gpt-4.1"] + assert models_route.calls.last.request.headers["Authorization"] == "Bearer discovery-bearer" + + @respx.mock + def test_discovers_with_env_wif_when_params_carry_no_key(self, wif_env: OpenAIWorkloadIdentityConfig) -> None: + mock_token_exchange("env-discovery-bearer") + models_route: Final = self.mock_models() + + OpenAIGPTConfig().discover_models({}) + + assert models_route.calls.last.request.headers["Authorization"] == "Bearer env-discovery-bearer" + + @respx.mock + def test_static_api_key_in_params_skips_token_exchange(self, deployment_wif: dict[str, str]) -> None: + exchange_route: Final = mock_token_exchange() + models_route: Final = self.mock_models() + + OpenAIGPTConfig().discover_models({**deployment_wif, "api_key": "sk-discovery"}) + + assert models_route.calls.last.request.headers["Authorization"] == "Bearer sk-discovery" + assert not exchange_route.called + + @respx.mock + def test_blank_api_base_in_params_discovers_from_openai(self, deployment_wif: dict[str, str]) -> None: + mock_token_exchange("blank-base-bearer") + models_route: Final = self.mock_models() + + assert OpenAIGPTConfig().discover_models({**deployment_wif, "api_base": ""}) == ["gpt-4o-mini", "gpt-4.1"] + assert models_route.calls.last.request.headers["Authorization"] == "Bearer blank-base-bearer" + + @respx.mock + def test_openai_compatible_subclass_never_mints_wif(self, deployment_wif: dict[str, str]) -> None: + exchange_route: Final = mock_token_exchange() + models_route: Final = self.mock_models() + + class CompatibleConfig(OpenAIGPTConfig): + pass + + CompatibleConfig().discover_models(deployment_wif) + + assert models_route.calls.last.request.headers["Authorization"] == "Bearer None" + assert not exchange_route.called + + + @respx.mock + def test_empty_static_key_never_borrows_the_env_key(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("OPENAI_API_KEY", "sk-env-key-that-must-stay-home") + foreign_models: Final = respx.get("https://third-party.example/v1/models").mock( + return_value=httpx.Response(200, json={"data": [{"id": "other-model"}]}) + ) + + assert OpenAIGPTConfig().get_models(api_key="", api_base="https://third-party.example") == ["other-model"] + assert foreign_models.calls.last.request.headers["Authorization"] == "Bearer " + + +class TestClientsideBaseOverride: + def test_client_api_base_override_clears_deployment_wif(self, deployment_wif: dict[str, str]) -> None: + from litellm.router_utils.clientside_credential_handler import get_dynamic_litellm_params + + redirected: Final = get_dynamic_litellm_params( + litellm_params={"model": "openai/gpt-4o-mini", **deployment_wif}, + request_kwargs={"api_base": "https://not-openai.example/v1"}, + ) + + assert not any(key in redirected for key in deployment_wif) + assert ( + resolve_openai_workload_identity_config( + api_key=None, api_base=redirected["api_base"], litellm_params=redirected + ) + is None + ) diff --git a/tests/unit/models/test_models.py b/tests/unit/models/test_models.py index 7b8953bd1a0..139df21decc 100644 --- a/tests/unit/models/test_models.py +++ b/tests/unit/models/test_models.py @@ -3,6 +3,7 @@ Tests for backend domain models. """ from datetime import datetime, timezone +from typing import Final import pytest from pydantic import BaseModel, TypeAdapter, ValidationError @@ -603,6 +604,43 @@ class TestManagedTables: assert table.custom_llm_provider == "openai" +class TestProxyModelTableResponseSerialization: + """FastAPI validates an endpoint's return value against its response model with + ``from_attributes``, so an endpoint that returns an already-built row reaches the + ``mode="before"`` validator as the object itself rather than as a mapping.""" + + def test_validates_from_an_existing_instance(self): + from pydantic import TypeAdapter + + built: Final = LiteLLM_ProxyModelTable( + model_id="m-1", + model_name="claude-sonnet-5-provider", + litellm_params={"model": "anthropic/claude-sonnet-5"}, + blocked=True, + ) + + serialized = TypeAdapter(LiteLLM_ProxyModelTable | None).validate_python(built, from_attributes=True) + + assert serialized is not None + assert serialized.model_id == "m-1" + assert serialized.blocked is True + assert serialized.litellm_params == {"model": "anthropic/claude-sonnet-5"} + + def test_still_parses_json_string_columns(self): + """The DB stores these columns as JSON strings, which is why the validator exists.""" + parsed: Final = LiteLLM_ProxyModelTable.model_validate( + { + "model_id": "m-2", + "model_name": "n", + "litellm_params": '{"model": "anthropic/claude-haiku-4-5"}', + "model_info": '{"id": "m-2"}', + } + ) + + assert parsed.litellm_params == {"model": "anthropic/claude-haiku-4-5"} + assert parsed.model_info == {"id": "m-2"} + + class TestAutoRouterSession: @staticmethod def _row(baseline_models: dict[str, int], estimated_turns: int = 3) -> LiteLLM_AutoRouterSession: diff --git a/tests/unit/proxy/auth/test_auth_utils.py b/tests/unit/proxy/auth/test_auth_utils.py index bc4a6e0155d..90e0595dc17 100644 --- a/tests/unit/proxy/auth/test_auth_utils.py +++ b/tests/unit/proxy/auth/test_auth_utils.py @@ -10,6 +10,7 @@ from unittest.mock import MagicMock, patch import pytest from fastapi import HTTPException, Request +import litellm from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.auth_utils import ( _get_customer_id_from_standard_headers, @@ -30,6 +31,181 @@ from litellm.proxy.auth.auth_utils import ( get_request_route_template, is_request_body_safe, ) +from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS, OPENAI_WIF_KWARGS_KEYS + + +@pytest.mark.parametrize("param", sorted(ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS)) +def test_every_wif_kwarg_key_is_refused_from_a_request_body(param: str): + """Every key the kwargs funnel carries into litellm_params selects a server-side secret or the + scope a token is minted for, so each one must be refused from a request body even with the + proxy-wide client-credential opt-in; a key added to the funnel without joining the ban shows up + here as a body the proxy accepted.""" + with pytest.raises(ValueError, match="server-owned workload identity federation parameter"): + is_request_body_safe( + request_body={"model": "claude-sonnet-5", param: "attacker-chosen"}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="claude-sonnet-5", + ) + + +@pytest.mark.parametrize( + "body", + [ + {"model": "claude-sonnet-5", "litellm_credential_name": "admin-wif"}, + {"model": "claude-sonnet-5", "litellm_params": {"litellm_credential_name": "admin-wif"}}, + ], + ids=["top_level", "nested_litellm_params"], +) +def test_a_request_body_cannot_pick_a_federated_identity_by_credential_name(monkeypatch, body: dict): + """Naming a federated credential moves the token exchange onto that credential's federation rule + and organization just as sending the fields inline does, so the ban on the inline form has to + cover the reference too.""" + from litellm.types.utils import CredentialItem + + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="admin-wif", + credential_values={ + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + with pytest.raises(ValueError, match="names a credential configured for workload identity federation"): + is_request_body_safe( + request_body=body, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="claude-sonnet-5", + ) + + +def test_a_request_body_may_still_name_a_credential_that_does_not_federate(monkeypatch): + """Only federation makes a credential a deployment decision. An ordinary named credential stays + usable from a request body, so the ban must read what the credential holds, not its presence.""" + from litellm.types.utils import CredentialItem + + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="plain-key", + credential_values={"api_key": "sk-plain"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + assert ( + is_request_body_safe( + request_body={"model": "claude-sonnet-5", "litellm_credential_name": "plain-key"}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="claude-sonnet-5", + ) + is True + ) + + +@pytest.fixture +def federated_credential(monkeypatch): + """A stored credential that federates, so a body naming it is the reference the ban targets.""" + from litellm.types.utils import CredentialItem + + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="admin-wif", + credential_values={ + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + +@pytest.mark.parametrize( + "route", + [ + "/model/new", + "/model/update", + "/model/delete", + "/model/f38d7ce5-7966-42f2-bd06-67ea74aeb76b/update", + "/health/test_connection", + ], +) +@pytest.mark.parametrize( + "body", + [ + {"model": "claude-sonnet-5", "litellm_credential_name": "admin-wif"}, + {"model": "claude-sonnet-5", "litellm_params": {"litellm_credential_name": "admin-wif"}}, + ], + ids=["top_level", "nested_litellm_params"], +) +def test_configuring_a_deployment_may_name_a_federated_credential(federated_credential, route: str, body: dict): + """Attaching a federated credential to a deployment is the decision the ban tells the caller to + make, and ModelManagementAuthChecks._reject_non_admin_wif_write is what judges it: it lets a + proxy admin through and refuses everyone else with a 403. Refusing the name here first would + leave no API or Admin UI path to configure federation at all.""" + assert ( + is_request_body_safe( + request_body=body, + general_settings={}, + llm_router=None, + model="claude-sonnet-5", + route=route, + ) + is True + ) + + +@pytest.mark.parametrize( + "route", + [ + None, + "/v1/chat/completions", + "/v1/messages", + "/model/info", + "/model/f38d7ce5-7966-42f2-bd06-67ea74aeb76b/update/extra", + ], +) +def test_a_call_still_cannot_pick_a_federated_identity_by_credential_name(federated_credential, route: str | None): + """The exemption covers the deployment-management routes and nothing that shares their prefix, + so a call still cannot move its token exchange onto a federated credential by naming it.""" + with pytest.raises(ValueError, match="names a credential configured for workload identity federation"): + is_request_body_safe( + request_body={"model": "claude-sonnet-5", "litellm_credential_name": "admin-wif"}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="claude-sonnet-5", + route=route, + ) + + +@pytest.mark.parametrize("route", ["/model/new", "/model/f38d7ce5-7966-42f2-bd06-67ea74aeb76b/update"]) +def test_configuring_a_deployment_still_cannot_carry_federation_fields_inline(route: str): + """Only the credential reference is exempt. Federation fields typed straight into a body stay + refused everywhere, since a stored credential is the surface an admin has to go through.""" + with pytest.raises(ValueError, match="server-owned workload identity federation parameter"): + is_request_body_safe( + request_body={"model": "claude-sonnet-5", "anthropic_federation_rule_id": "fdrl_attacker"}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="claude-sonnet-5", + route=route, + ) class TestCustomAuthCommonChecksWarning: @@ -185,9 +361,7 @@ class TestGetKeyModelRpmLimit: """Should fall back to team metadata when key metadata exists but has no model_rpm_limit.""" user_api_key_dict = UserAPIKeyAuth( api_key="sk-123", - metadata={ - "some_other_key": "value" - }, # Has metadata, but not model_rpm_limit + metadata={"some_other_key": "value"}, # Has metadata, but not model_rpm_limit team_metadata={"model_rpm_limit": {"gpt-4": 50}}, ) result = get_key_model_rpm_limit(user_api_key_dict) @@ -269,9 +443,7 @@ class TestGetKeyModelTpmLimit: """Should fall back to team metadata when key metadata exists but has no model_tpm_limit.""" user_api_key_dict = UserAPIKeyAuth( api_key="sk-123", - metadata={ - "some_other_key": "value" - }, # Has metadata, but not model_tpm_limit + metadata={"some_other_key": "value"}, # Has metadata, but not model_tpm_limit team_metadata={"model_tpm_limit": {"gpt-4": 5000}}, ) result = get_key_model_tpm_limit(user_api_key_dict) @@ -382,9 +554,7 @@ class TestGetEndUserIdFromRequestBodyWithStandardHeaders: request_body = {"user": "body-user"} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers) assert result == "header-customer" def test_should_fall_back_to_body_when_no_standard_header(self): @@ -393,9 +563,7 @@ class TestGetEndUserIdFromRequestBodyWithStandardHeaders: request_body = {"user": "body-user"} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers) assert result == "body-user" @@ -437,8 +605,7 @@ def test_get_model_from_request_enforces_when_builtin_handler_dispatched(): enforced. Same request path as above, but dispatched to a non-pass-through endpoint: the model must NOT be suppressed.""" - def builtin_chat_completions(): - ... + def builtin_chat_completions(): ... assert ( get_model_from_request( @@ -1015,9 +1182,7 @@ def test_get_model_from_request_extracts_unified_file_id_models(): "litellm_proxy:application/octet-stream;unified_id,test-id;" "target_model_names,model-a,model-b;llm_output_file_id,file-provider-id" ) - encoded_unified_file_id = ( - base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=") - ) + encoded_unified_file_id = base64.urlsafe_b64encode(raw_unified_file_id.encode()).decode().rstrip("=") assert get_model_from_request( request_data={"file_id": encoded_unified_file_id}, @@ -1077,9 +1242,7 @@ def test_get_model_from_request_resolves_video_id_model_with_router(): model_id="veo-3.1-generate-001", ) llm_router = MagicMock() - llm_router.resolve_model_name_from_model_id.return_value = ( - "gcp/google/veo-3.1-generate-001" - ) + llm_router.resolve_model_name_from_model_id.return_value = "gcp/google/veo-3.1-generate-001" assert ( get_model_from_request( @@ -1089,9 +1252,7 @@ def test_get_model_from_request_resolves_video_id_model_with_router(): ) == "gcp/google/veo-3.1-generate-001" ) - llm_router.resolve_model_name_from_model_id.assert_called_once_with( - "veo-3.1-generate-001" - ) + llm_router.resolve_model_name_from_model_id.assert_called_once_with("veo-3.1-generate-001") _BATCH_DEPLOYMENT_ID = "8d0eaa7e6c6f54a425dfd0062cb6b0dc" @@ -1122,9 +1283,7 @@ def _encode_managed_id(decoded: str) -> str: return base64.urlsafe_b64encode(decoded.encode()).decode().rstrip("=") -_MANAGED_BATCH_ID = _encode_managed_id( - f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123" -) +_MANAGED_BATCH_ID = _encode_managed_id(f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123") _MANAGED_BATCH_OUTPUT_FILE_ID = _encode_managed_id( f"litellm_proxy;model_id:{_BATCH_DEPLOYMENT_ID};llm_batch_id:provider-batch-123;" "llm_output_file_id:provider-file-456" @@ -1200,9 +1359,7 @@ def test_get_model_from_request_resolves_character_id_model_with_router(): model_id="veo-3.1-generate-001", ) llm_router = MagicMock() - llm_router.resolve_model_name_from_model_id.return_value = ( - "gcp/google/veo-3.1-generate-001" - ) + llm_router.resolve_model_name_from_model_id.return_value = "gcp/google/veo-3.1-generate-001" assert ( get_model_from_request( @@ -1212,9 +1369,7 @@ def test_get_model_from_request_resolves_character_id_model_with_router(): ) == "gcp/google/veo-3.1-generate-001" ) - llm_router.resolve_model_name_from_model_id.assert_called_once_with( - "veo-3.1-generate-001" - ) + llm_router.resolve_model_name_from_model_id.assert_called_once_with("veo-3.1-generate-001") def test_get_model_from_request_only_runs_media_decoders_for_matching_fields(): @@ -1361,9 +1516,7 @@ def test_abbreviate_api_key_short_key_is_fully_masked(): def test_get_customer_user_header_returns_none_when_no_customer_role(): from litellm.proxy.auth.auth_utils import get_customer_user_header_from_mapping - mappings = [ - {"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"} - ] + mappings = [{"header_name": "X-OpenWebUI-User-Id", "litellm_user_role": "internal_user"}] result = get_customer_user_header_from_mapping(mappings) assert result is None @@ -1416,9 +1569,7 @@ def test_get_end_user_id_returns_id_from_user_header_mappings(): ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body={}, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result == "1234" @@ -1444,9 +1595,7 @@ def test_get_end_user_id_returns_first_customer_header_when_multiple_mappings_ex ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body={}, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result == "user-456" @@ -1467,9 +1616,7 @@ def test_get_end_user_id_returns_none_when_no_customer_role_in_mappings(): ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body={}, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result is None @@ -1487,9 +1634,7 @@ def test_get_end_user_id_falls_back_to_deprecated_user_header_name(): ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body={}, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body={}, request_headers=headers) assert result == "user-legacy" @@ -1633,9 +1778,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "alice@example.com" @@ -1645,9 +1788,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result is None @@ -1660,19 +1801,14 @@ class TestGetEndUserIdDropsMalformedBodyValues: """ import litellm - blob = ( - '{"device_id":"d5abe9199ee7759a","account_uuid":"",' - '"session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}' - ) + blob = '{"device_id":"d5abe9199ee7759a","account_uuid":"","session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}' request_body = {"user": blob} original = litellm.validate_end_user_id_in_db litellm.validate_end_user_id_in_db = False try: with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) finally: litellm.validate_end_user_id_in_db = original @@ -1683,8 +1819,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = { "user": ( - '{"device_id":"d5abe9199ee7759a","account_uuid":"",' - '"session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}' + '{"device_id":"d5abe9199ee7759a","account_uuid":"","session_id":"c284b8cb-a050-4278-8599-cc4e016a10ab"}' ), } @@ -1692,9 +1827,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: litellm.validate_end_user_id_in_db = True try: with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) finally: litellm.validate_end_user_id_in_db = original @@ -1704,9 +1837,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = {"user": "alice@example.com"} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "alice@example.com" @@ -1718,9 +1849,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = {"user": codex_id} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == codex_id @@ -1728,9 +1857,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = {"user": 12345} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "12345" @@ -1741,9 +1868,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "alice@example.com" @@ -1753,9 +1878,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result is None @@ -1765,9 +1888,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: } with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result is None @@ -1775,9 +1896,7 @@ class TestGetEndUserIdDropsMalformedBodyValues: request_body = {"user": " ", "safety_identifier": "alice@example.com"} with patch("litellm.proxy.proxy_server.general_settings", {}): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers={} - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers={}) assert result == "alice@example.com" @@ -1796,16 +1915,12 @@ class TestGetEndUserIdDropsMalformedBodyValues: ), patch("litellm.proxy.proxy_server.general_settings", general_settings), ): - result = get_end_user_id_from_request_body( - request_body=request_body, request_headers=headers - ) + result = get_end_user_id_from_request_body(request_body=request_body, request_headers=headers) assert result == "alice@example.com" -def _make_deployment_dict( - model_name: str, tpm: Optional[int] = None, rpm: Optional[int] = None -) -> dict: +def _make_deployment_dict(model_name: str, tpm: Optional[int] = None, rpm: Optional[int] = None) -> dict: """Helper to build a minimal deployment dict as returned by router.get_model_list.""" litellm_params: dict = {"model": model_name} if tpm is not None: @@ -1825,9 +1940,7 @@ class TestDeploymentDefaultRpmLimit: """Case 2 from spec: key has no model-specific limits, falls back to deployment default.""" user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", rpm=200) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") assert result == {"model1": 200} @@ -1839,9 +1952,7 @@ class TestDeploymentDefaultRpmLimit: metadata={"model_rpm_limit": {"model1": 10}}, ) mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", rpm=200) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_rpm_limit(user_api_key_dict, model_name="model1") assert result == {"model1": 10} @@ -1861,9 +1972,7 @@ class TestDeploymentDefaultRpmLimit: """No model_name means deployment fallback is skipped.""" user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", rpm=200) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", rpm=200)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_rpm_limit(user_api_key_dict) assert result is None @@ -1924,9 +2033,7 @@ class TestDeploymentDefaultTpmLimit: """Case 2 from spec: key has no model-specific limits, falls back to deployment default.""" user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", tpm=100) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") assert result == {"model1": 100} @@ -1938,9 +2045,7 @@ class TestDeploymentDefaultTpmLimit: metadata={"model_tpm_limit": {"model1": 20}}, ) mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", tpm=100) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_tpm_limit(user_api_key_dict, model_name="model1") assert result == {"model1": 20} @@ -1960,9 +2065,7 @@ class TestDeploymentDefaultTpmLimit: """No model_name means deployment fallback is skipped.""" user_api_key_dict = UserAPIKeyAuth(api_key="sk-123") mock_router = MagicMock() - mock_router.get_model_list.return_value = [ - _make_deployment_dict("model1", tpm=100) - ] + mock_router.get_model_list.return_value = [_make_deployment_dict("model1", tpm=100)] with patch(_ROUTER_PATCH, mock_router): result = get_key_model_tpm_limit(user_api_key_dict) assert result is None @@ -2113,7 +2216,7 @@ class TestCheckCompleteCredentialsBlocksSSRF: "litellm.proxy.auth.auth_utils.validate_url", side_effect=SSRFError(f"blocked: {blocked_url}"), ): - with pytest.raises(ValueError, match='is rejected by the SSRF guard') as exc_info: + with pytest.raises(ValueError, match="is rejected by the SSRF guard") as exc_info: check_complete_credentials( { "model": "gpt-4", @@ -2475,9 +2578,7 @@ class TestIsRequestBodySafeBlocksFallbackSmuggle: is_request_body_safe( request_body={ "model": "gpt-4", - "fallbacks": [ - {"gpt-4": [{"model": "byok", "api_base": "https://my-byok.example"}]} - ], + "fallbacks": [{"gpt-4": [{"model": "byok", "api_base": "https://my-byok.example"}]}], }, general_settings={"allow_client_side_credentials": True}, llm_router=None, @@ -2497,9 +2598,7 @@ class TestIsRequestBodySafeBlocksFallbackSmuggle: "always-fail": [ { "model": "x", - fallback_field: [ - {"x": [{"model": "deepseek-chat", "api_base": "http://attacker"}]} - ], + fallback_field: [{"x": [{"model": "deepseek-chat", "api_base": "http://attacker"}]}], } ] } @@ -2669,7 +2768,7 @@ class TestIsRequestBodySafeBlocksEndpointTargetingFields: ], ) def test_endpoint_targeting_field_in_request_body_is_rejected(self, field): - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={"model": "gpt-4", field: "https://attacker.example"}, general_settings={}, @@ -2690,7 +2789,7 @@ class TestIsRequestBodySafeBlocksEndpointTargetingFields: # on the blocklist into an SSRF / credential-exfil hole. Verify # that supplying an api_key (alongside the banned param) does NOT # bypass the gate — it can only be opened by an admin opt-in. - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={ "model": "gpt-4", @@ -3139,11 +3238,7 @@ class TestIsRequestBodySafeNestedConfig: when nested.""" with pytest.raises(ValueError, match="langfuse_host"): is_request_body_safe( - request_body={ - "litellm_embedding_config": { - "langfuse_host": "https://attacker.example.com" - } - }, + request_body={"litellm_embedding_config": {"langfuse_host": "https://attacker.example.com"}}, general_settings={}, llm_router=None, model="milvus-store", @@ -3154,11 +3249,7 @@ class TestIsRequestBodySafeNestedConfig: keep the existing escape hatch — same UX as for root-level.""" assert ( is_request_body_safe( - request_body={ - "litellm_embedding_config": { - "api_base": "https://my-azure.example.com" - } - }, + request_body={"litellm_embedding_config": {"api_base": "https://my-azure.example.com"}}, general_settings={"allow_client_side_credentials": True}, llm_router=None, model="milvus-store", @@ -3271,7 +3362,7 @@ class TestObservabilityCallbackBans: ], ) def test_observability_field_in_request_body_root_is_rejected(self, field): - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={"model": "gpt-4", field: "attacker-value"}, general_settings={}, @@ -3295,13 +3386,11 @@ class TestObservabilityCallbackBans: "user_api_key_auth_metadata", ], ) - def test_observability_field_in_metadata_dict_is_rejected( - self, metadata_key, field - ): + def test_observability_field_in_metadata_dict_is_rejected(self, metadata_key, field): # Verifies the metadata walk: a value smuggled inside ``metadata`` # or ``litellm_metadata`` is just as dangerous as the same field # at the body root, and must hit the same gate. - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={ "model": "gpt-4", @@ -3336,13 +3425,11 @@ class TestObservabilityCallbackBans: ) def test_observability_field_in_litellm_params_metadata_is_rejected(self): - with pytest.raises(ValueError, match='Rejected Request: turn_off_message_logging is not allowed') as exc: + with pytest.raises(ValueError, match="Rejected Request: turn_off_message_logging is not allowed") as exc: is_request_body_safe( request_body={ "model": "gpt-4", - "litellm_params": { - "metadata": {"turn_off_message_logging": False} - }, + "litellm_params": {"metadata": {"turn_off_message_logging": False}}, }, general_settings={}, llm_router=None, @@ -3354,22 +3441,18 @@ class TestObservabilityCallbackBans: "metadata_key", ["metadata", "litellm_metadata"], ) - def test_observability_field_in_json_string_metadata_is_rejected( - self, metadata_key - ): + def test_observability_field_in_json_string_metadata_is_rejected(self, metadata_key): # Multipart/form-data and ``extra_body`` callers send metadata as a # JSON-encoded string. The bouncer parses it before applying the # banned-params check so the JSON-string path can't smuggle past # the ``isinstance(dict)`` guard. import json - with pytest.raises(ValueError, match='Rejected Request: langfuse_host is not allowed in request') as exc: + with pytest.raises(ValueError, match="Rejected Request: langfuse_host is not allowed in request") as exc: is_request_body_safe( request_body={ "model": "gpt-4", - metadata_key: json.dumps( - {"langfuse_host": "https://attacker.example"} - ), + metadata_key: json.dumps({"langfuse_host": "https://attacker.example"}), }, general_settings={}, llm_router=None, @@ -3436,7 +3519,7 @@ def test_model_level_allow_does_not_skip_subsequent_banned_params(monkeypatch): lambda model, param, request_body_value, llm_router: param == "api_base", ) - with pytest.raises(ValueError, match='Rejected Request: langfuse_host is not allowed in request') as exc: + with pytest.raises(ValueError, match="Rejected Request: langfuse_host is not allowed in request") as exc: is_request_body_safe( request_body={ "model": "gpt-4", @@ -3476,8 +3559,7 @@ def test_observability_ban_covers_canonical_supported_callback_params(): ) for param in _request_blocked_callback_params: assert param in banned, ( - f"{param} is in _request_blocked_callback_params but is not banned " - "at the proxy request-body boundary." + f"{param} is in _request_blocked_callback_params but is not banned at the proxy request-body boundary." ) @@ -3507,7 +3589,7 @@ class TestPricingInjectionBlocked: ], ) def test_pricing_field_rejected_by_default(self, field, value): - with pytest.raises(ValueError, match='Rejected Request') as exc: + with pytest.raises(ValueError, match="Rejected Request") as exc: is_request_body_safe( request_body={"model": "gpt-4", field: value}, general_settings={}, @@ -3575,9 +3657,7 @@ class TestGetRequestRouteTemplate: def test_exception_returns_none(self): req = MagicMock() - type(req).scope = property( - lambda self: (_ for _ in ()).throw(RuntimeError("boom")) - ) + type(req).scope = property(lambda self: (_ for _ in ()).throw(RuntimeError("boom"))) assert get_request_route_template(req) is None @@ -3630,9 +3710,7 @@ class TestGetKeyTagRateLimits: """Tests for get_key_tag_rpm_limit.""" def test_reads_tag_rpm_limit_from_metadata(self): - key = UserAPIKeyAuth( - api_key="sk-123", metadata={"tag_rpm_limit": {"cell-1": 5}} - ) + key = UserAPIKeyAuth(api_key="sk-123", metadata={"tag_rpm_limit": {"cell-1": 5}}) assert get_key_tag_rpm_limit(key) == {"cell-1": 5} def test_returns_none_when_unset(self): @@ -3695,12 +3773,8 @@ class TestIsRequestBodySafeChecksBracketNotationMetadata: def test_bracket_notation_matches_json_encoding_for_deeper_nesting(self): """A value nested below the first level is treated the same either way: the check descends one level into metadata, for both encodings.""" - deep_bracket = { - "litellm_metadata[spend_logs_metadata][langfuse_host]": "https://example.invalid" - } - deep_json = { - "litellm_metadata": {"spend_logs_metadata": {"langfuse_host": "https://example.invalid"}} - } + deep_bracket = {"litellm_metadata[spend_logs_metadata][langfuse_host]": "https://example.invalid"} + deep_json = {"litellm_metadata": {"spend_logs_metadata": {"langfuse_host": "https://example.invalid"}}} kwargs = dict(general_settings={}, llm_router=None, model="gpt-4") assert is_request_body_safe(request_body=deep_bracket, **kwargs) is True assert is_request_body_safe(request_body=deep_json, **kwargs) is True @@ -3749,9 +3823,7 @@ class TestHasUserSetupSso: def test_true_for_saml_metadata_url(self, monkeypatch): from litellm.proxy.auth.auth_utils import has_user_setup_sso - monkeypatch.setenv( - "SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml" - ) + monkeypatch.setenv("SAML_IDP_METADATA_URL", "https://idp.example.com/metadata.xml") assert has_user_setup_sso() is True def test_true_for_saml_metadata_xml(self, monkeypatch): diff --git a/tests/unit/proxy/common_utils/test_credential_hydration.py b/tests/unit/proxy/common_utils/test_credential_hydration.py new file mode 100644 index 00000000000..f40b0114d71 --- /dev/null +++ b/tests/unit/proxy/common_utils/test_credential_hydration.py @@ -0,0 +1,28 @@ +from unittest.mock import AsyncMock, MagicMock, patch + +import pytest + +import litellm +from litellm.proxy.common_utils.credential_hydration import hydrate_named_credential_authoritative +from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + +@pytest.mark.asyncio +async def test_authoritative_hydrate_returns_an_encrypted_empty_value_as_empty(monkeypatch): + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-hydration-test-salt") + row = { + "credential_name": "openai-wif", + "credential_values": { + "api_base": encrypt_value_helper(""), + "openai_service_account_id": encrypt_value_helper("user-1"), + }, + "credential_info": {"custom_llm_provider": "openai"}, + } + prisma = MagicMock() + prisma.db.litellm_credentialstable.find_unique = AsyncMock(return_value=row) + + with patch.object(litellm, "credential_list", []): # test-quality-ok: the row under test must win over memory + resolved = await hydrate_named_credential_authoritative("openai-wif", prisma) + + assert resolved is not None + assert resolved.credential_values == {"api_base": "", "openai_service_account_id": "user-1"} diff --git a/tests/unit/proxy/credential_endpoints/test_endpoints.py b/tests/unit/proxy/credential_endpoints/test_endpoints.py index 631767f52ae..aff7babe424 100644 --- a/tests/unit/proxy/credential_endpoints/test_endpoints.py +++ b/tests/unit/proxy/credential_endpoints/test_endpoints.py @@ -1,12 +1,14 @@ """Tests for the credential management endpoints.""" import json +from contextlib import contextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import ec from fastapi.testclient import TestClient - import litellm from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth @@ -21,10 +23,14 @@ def _as_admin(): return UserAPIKeyAuth(api_key="test-key", user_role="proxy_admin") -def _call_as_admin(method: str, path: str, json_body: dict | None = None): +def _as_non_admin(): + return UserAPIKeyAuth(api_key="test-key", user_role="internal_user") + + +def _call_as(method: str, path: str, json_body: dict | None = None, auth=_as_admin): missing = object() previous_override = app.dependency_overrides.get(user_api_key_auth, missing) - app.dependency_overrides[user_api_key_auth] = _as_admin + app.dependency_overrides[user_api_key_auth] = auth try: return client.request(method, path, json=json_body, headers={"Authorization": "Bearer test-key"}) finally: @@ -34,16 +40,26 @@ def _call_as_admin(method: str, path: str, json_body: dict | None = None): app.dependency_overrides[user_api_key_auth] = previous_override -def _patch_credential(name: str, body: dict): - return _call_as_admin("PATCH", f"/credentials/{name}", body) +def _patch_credential(name: str, body: dict, auth=_as_admin): + return _call_as("PATCH", f"/credentials/{name}", body, auth) -def _delete_credential(name: str): - return _call_as_admin("DELETE", f"/credentials/{name}") +def _post_credential(body: dict, auth=_as_admin): + return _call_as("POST", "/credentials", body, auth) + + +def _delete_credential(name: str, auth=_as_admin): + return _call_as("DELETE", f"/credentials/{name}", auth=auth) def _list_credentials(): - return _call_as_admin("GET", "/credentials") + return _call_as("GET", "/credentials") + + +def _prisma_without_credential_rows() -> MagicMock: + prisma_client = MagicMock() + prisma_client.db.litellm_credentialstable.find_unique = AsyncMock(return_value=None) + return prisma_client @pytest.fixture @@ -59,11 +75,12 @@ def credential_store(): llm_router: object | None = None, **repository_calls: AsyncMock, ) -> None: - patch("litellm.proxy.proxy_server.prisma_client", MagicMock() if connected else None).start() + patch("litellm.proxy.proxy_server.prisma_client", _prisma_without_credential_rows() if connected else None).start() patch("litellm.proxy.proxy_server.master_key", "sk-test-master").start() patch.object(litellm, "credential_list", list(in_memory)).start() app.dependency_overrides[get_llm_router] = lambda: llm_router repository = patch("litellm.proxy.credential_endpoints.endpoints.CredentialsRepository").start() + repository.return_value.find_by_name = AsyncMock(return_value=None) for call_name, result in repository_calls.items(): setattr(repository.return_value, call_name, result) @@ -72,6 +89,52 @@ def credential_store(): app.dependency_overrides.pop(get_llm_router, None) +@contextmanager +def _repository_holding(stored: CredentialItem | None): + """The credentials repository seam, answering ``find_by_name`` with ``stored`` and recording + the writes the handler attempts. Patched at both import sites, since the handlers resolve an + existing credential through ``hydrate_named_credential`` (memory first, then this repository) + and then write through their own ``CredentialsRepository`` binding.""" + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.credential_endpoints.endpoints.CredentialsRepository" + ) as repository, + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.common_utils.credential_hydration.CredentialsRepository", repository + ), + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + repository.return_value.create = AsyncMock(return_value=None) + repository.return_value.update_by_name = AsyncMock(return_value=None) + repository.return_value.delete_by_name = AsyncMock(return_value=stored) + yield repository.return_value + + +def test_create_credential_write_omits_the_patch_only_deletion_field(restore_credential_list): + """Regression: CredentialItem.credential_values_to_delete is a PATCH-only field that + defaults to None on every other construction path. A bare .model_dump() (without + exclude_none) on the create path put a `credential_values_to_delete: null` key into the + Prisma write, which litellm_credentialstable has no column for.""" + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "new-cred", + "credential_values": {"api_key": "sk-new"}, + "credential_info": {"custom_llm_provider": "openai"}, + } + ) + + assert response.status_code == 200, response.text + written_data = repository.create.await_args.kwargs["data"] + assert "credential_values_to_delete" not in written_data + + def test_update_credential_answers_404_when_the_credential_does_not_exist(credential_store): """Regression: the handler used to ``return handle_exception_on_proxy(e)``, which makes the exception the response body and lets FastAPI answer 200, so a write the handler @@ -119,6 +182,1012 @@ def test_update_credential_still_answers_200_on_a_successful_write(credential_st assert response.json()["success"] is True +def _get_jwks(name: str): + return _call_as("GET", f"/credentials/{name}/jwks") + + +@pytest.fixture +def restore_credential_list(monkeypatch): + monkeypatch.setattr(litellm, "credential_list", []) + + +def test_update_credential_rejects_overlap_between_update_and_delete(): + """A key in both sets is ambiguous (set to what value, before or after the delete?), so the + endpoint must reject it outright rather than picking a resolution order silently.""" + response = _patch_credential( + "any-name", + { + "credential_name": "any-name", + "credential_values": {"api_key": "sk-new"}, + "credential_values_to_delete": ["api_key"], + "credential_info": {}, + }, + ) + + assert response.status_code == 400, response.text + assert "api_key" in response.json()["error"]["message"] + + +def test_update_credential_deletion_removes_the_key_from_the_db_write(restore_credential_list): + """The bug this closes: switching WIF identity sources (or WIF -> api_key) left the old + variant's fields behind in the DB row, which wif.py then rejects by presence.""" + stored = CredentialItem( + credential_name="wif-cred", + credential_values={"anthropic_identity_source": "keycloak", "anthropic_keycloak_client_id": "old-client"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.credential_endpoints.endpoints.CredentialsRepository" + ) as repository, # test-quality-ok: the proxy wiring under test is what this patches + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + update_mock = AsyncMock(return_value=None) + repository.return_value.update_by_name = update_mock + + response = _patch_credential( + "wif-cred", + { + "credential_name": "wif-cred", + "credential_values": {}, + "credential_values_to_delete": ["anthropic_keycloak_client_id"], + "credential_info": {}, + }, + ) + + assert response.status_code == 200, response.text + written_values = json.loads(update_mock.await_args.kwargs["data"]["credential_values"]) + assert "anthropic_keycloak_client_id" not in written_values + assert written_values["anthropic_identity_source"] == "keycloak" + + +def test_update_credential_deletion_updates_in_memory_credential_list(restore_credential_list, monkeypatch): + """The in-memory list is what the request-time auth resolvers read; a deletion that only + landed in the DB would leave the stale field servable until the next process restart.""" + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="wif-cred", + credential_values={ + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_client_id": "old-client", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + stored = CredentialItem( + credential_name="wif-cred", + credential_values={"anthropic_identity_source": "keycloak", "anthropic_keycloak_client_id": "old-client"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.credential_endpoints.endpoints.CredentialsRepository" + ) as repository, # test-quality-ok: the proxy wiring under test is what this patches + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + repository.return_value.update_by_name = AsyncMock(return_value=None) + + response = _patch_credential( + "wif-cred", + { + "credential_name": "wif-cred", + "credential_values": {}, + "credential_values_to_delete": ["anthropic_keycloak_client_id"], + "credential_info": {}, + }, + ) + + assert response.status_code == 200, response.text + in_memory = next(c for c in litellm.credential_list if c.credential_name == "wif-cred") + assert "anthropic_keycloak_client_id" not in in_memory.credential_values + assert in_memory.credential_values["anthropic_identity_source"] == "keycloak" + + +def test_update_credential_leaves_untouched_fields_alone(): + """Regression for the masked-value hazard: GET /credentials masks values, so a PATCH that + only names the field being changed must not let an untouched field be nulled or overwritten + by anything a round-tripped (masked) form value could contain.""" + stored = CredentialItem( + credential_name="existing", + credential_values={"api_key": "sk-real-value", "api_base": "https://api.anthropic.com"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.credential_endpoints.endpoints.CredentialsRepository" + ) as repository, # test-quality-ok: the proxy wiring under test is what this patches + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + update_mock = AsyncMock(return_value=None) + repository.return_value.update_by_name = update_mock + + response = _patch_credential( + "existing", + {"credential_name": "existing", "credential_values": {"api_key": "sk-rotated"}, "credential_info": {}}, + ) + + assert response.status_code == 200, response.text + written_values = json.loads(update_mock.await_args.kwargs["data"]["credential_values"]) + assert written_values["api_base"] == "https://api.anthropic.com" + + +def test_create_credential_never_stores_a_null_credential_value(restore_credential_list): + """The dashboard posts a key for every field on the provider's form, and the ones the operator + left blank arrive as null. A null carries no credential, and the federation resolver refuses a + foreign variant's field by key, so a stored null wedges every deployment naming this credential.""" + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "new-cred", + "credential_values": {"api_key": "sk-new", "anthropic_issuer_url": None}, + "credential_info": {"custom_llm_provider": "anthropic"}, + } + ) + + assert response.status_code == 200, response.text + written_values = json.loads(repository.create.await_args.kwargs["data"]["credential_values"]) + assert "anthropic_issuer_url" not in written_values + assert "api_key" in written_values + + +def test_update_credential_never_stores_a_null_credential_value(restore_credential_list): + """Same null on the update path, where the merge writes the whole row back: the field the null + named keeps whatever it stored, since removing a field is what credential_values_to_delete is for.""" + stored = CredentialItem( + credential_name="wif-cred", + credential_values={"anthropic_identity_source": "keycloak", "anthropic_keycloak_client_id": "old-client"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with _repository_holding(stored) as repository: + response = _patch_credential( + "wif-cred", + { + "credential_name": "wif-cred", + "credential_values": {"anthropic_keycloak_client_id": None}, + "credential_info": {}, + }, + ) + + assert response.status_code == 200, response.text + written_values = json.loads(repository.update_by_name.await_args.kwargs["data"]["credential_values"]) + assert written_values["anthropic_keycloak_client_id"] == "old-client" + + +def test_update_credential_never_syncs_a_null_into_the_in_memory_credential(restore_credential_list, monkeypatch): + """The in-memory list is what request-time resolution reads, so a null that only got kept out of + the DB row would still wedge every deployment until the next restart.""" + in_memory = CredentialItem( + credential_name="plain-cred", + credential_values={"api_key": "sk-old"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + monkeypatch.setattr(litellm, "credential_list", [in_memory]) + with _repository_holding( + CredentialItem( + credential_name="plain-cred", + credential_values={"api_key": "sk-old"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ): + response = _patch_credential( + "plain-cred", + { + "credential_name": "plain-cred", + "credential_values": {"api_key": "sk-rotated", "anthropic_issuer_url": None}, + "credential_info": {}, + }, + ) + + assert response.status_code == 200, response.text + synced = next(c for c in litellm.credential_list if c.credential_name == "plain-cred") + assert "anthropic_issuer_url" not in synced.credential_values + assert synced.credential_values["api_key"] == "sk-rotated" + + +def _generate_es256_pem() -> str: + key = ec.generate_private_key(ec.SECP256R1()) + return key.private_bytes( + encoding=serialization.Encoding.PEM, + format=serialization.PrivateFormat.PKCS8, + encryption_algorithm=serialization.NoEncryption(), + ).decode() + + +class TestCredentialJwksExport: + def test_jwks_export_succeeds_for_an_internal_issuer_credential(self, restore_credential_list, monkeypatch): + monkeypatch.setenv("JWKS_TEST_SIGNING_KEY", _generate_es256_pem()) + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="anthropic-issuer", + credential_values={ + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.example.com", + "anthropic_issuer_subject": "my-workload", + "anthropic_issuer_signing_key_ref": "os.environ/JWKS_TEST_SIGNING_KEY", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + response = _get_jwks("anthropic-issuer") + + assert response.status_code == 200, response.text + body = response.json() + assert body["keys"][0]["kty"] == "EC" + assert body["keys"][0]["crv"] == "P-256" + # The private key material must never leave the process via this endpoint. + assert "JWKS_TEST_SIGNING_KEY" not in response.text + assert "PRIVATE KEY" not in response.text + + def test_jwks_export_treats_blank_optional_fields_as_unset(self, restore_credential_list, monkeypatch): + monkeypatch.setenv("JWKS_TEST_SIGNING_KEY", _generate_es256_pem()) + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="anthropic-issuer-blanks", + credential_values={ + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.example.com", + "anthropic_issuer_subject": "my-workload", + "anthropic_issuer_signing_key_ref": "os.environ/JWKS_TEST_SIGNING_KEY", + "anthropic_issuer_audience": "", + "anthropic_issuer_ttl_seconds": "", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + response = _get_jwks("anthropic-issuer-blanks") + + assert response.status_code == 200, response.text + assert response.json()["keys"][0]["kty"] == "EC" + + def test_jwks_export_accepts_the_dashboard_provider_casing(self, restore_credential_list, monkeypatch): + monkeypatch.setenv("JWKS_TEST_SIGNING_KEY", _generate_es256_pem()) + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="anthropic-from-modal", + credential_values={ + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.example.com", + "anthropic_issuer_subject": "my-workload", + "anthropic_issuer_signing_key_ref": "os.environ/JWKS_TEST_SIGNING_KEY", + }, + credential_info={"custom_llm_provider": "Anthropic"}, + ) + ], + ) + + response = _get_jwks("anthropic-from-modal") + + assert response.status_code == 200, response.text + assert response.json()["keys"][0]["kty"] == "EC" + + def test_jwks_export_404s_for_a_non_anthropic_credential(self, restore_credential_list, monkeypatch): + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="openai-key", + credential_values={"api_key": "sk-x"}, + credential_info={"custom_llm_provider": "openai"}, + ) + ], + ) + + response = _get_jwks("openai-key") + + assert response.status_code == 404, response.text + + def test_jwks_export_404s_for_an_anthropic_credential_without_internal_issuer( + self, restore_credential_list, monkeypatch + ): + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="anthropic-apikey", + credential_values={"api_key": "sk-ant"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + response = _get_jwks("anthropic-apikey") + + assert response.status_code == 404, response.text + + def test_jwks_export_404s_for_an_unknown_credential(self, restore_credential_list): + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", None + ): # test-quality-ok: the proxy wiring under test is what this patches + response = _get_jwks("does-not-exist") + + assert response.status_code == 404, response.text + + def test_jwks_export_requires_proxy_admin(self, restore_credential_list, monkeypatch): + monkeypatch.setenv("JWKS_TEST_SIGNING_KEY", _generate_es256_pem()) + monkeypatch.setattr( + litellm, + "credential_list", + [ + CredentialItem( + credential_name="anthropic-issuer", + credential_values={ + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://issuer.example.com", + "anthropic_issuer_subject": "my-workload", + "anthropic_issuer_signing_key_ref": "os.environ/JWKS_TEST_SIGNING_KEY", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + ], + ) + + def _as_internal_user(): + return UserAPIKeyAuth(api_key="test-key", user_role="internal_user") + + app.dependency_overrides[user_api_key_auth] = _as_internal_user + try: + response = client.get("/credentials/anthropic-issuer/jwks", headers={"Authorization": "Bearer test-key"}) + finally: + app.dependency_overrides.pop(user_api_key_auth, None) + + assert response.status_code == 403, response.text + + +class TestNonAdminCannotPersistWifFieldsOnCredential: + """A credential's ``credential_values`` feeds the same WIF resolution as a deployment's own + ``litellm_params`` when referenced by ``litellm_credential_name``. A non-admin must not be + able to create or update a credential carrying a server-owned WIF field such as + ``anthropic_keycloak_token_url`` (destination) or ``anthropic_keycloak_client_secret_ref`` + (which secret to read and send there).""" + + def test_non_admin_cannot_create_a_credential_with_a_wif_destination(self): + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ): # test-quality-ok: the proxy wiring under test is what this patches + response = _post_credential( + { + "credential_name": "attacker-cred", + "credential_values": {"anthropic_keycloak_token_url": "https://evil.example.com/token"}, + "credential_info": {"custom_llm_provider": "anthropic"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "anthropic_keycloak_token_url" in response.json()["error"]["message"] + + def test_non_admin_cannot_create_a_credential_with_a_wif_secret_ref(self): + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ): # test-quality-ok: the proxy wiring under test is what this patches + response = _post_credential( + { + "credential_name": "attacker-cred", + "credential_values": {"anthropic_keycloak_client_secret_ref": "os.environ/LITELLM_MASTER_KEY"}, + "credential_info": {"custom_llm_provider": "anthropic"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + + def test_non_admin_can_create_a_credential_without_wif_fields(self, restore_credential_list): + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "ordinary-cred", + "credential_values": {"api_key": "sk-new"}, + "credential_info": {"custom_llm_provider": "openai"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 200, response.text + repository.create.assert_awaited_once() + + def test_proxy_admin_can_create_a_credential_with_a_wif_destination(self, restore_credential_list): + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "admin-cred", + "credential_values": {"anthropic_keycloak_token_url": "https://keycloak.internal/token"}, + "credential_info": {"custom_llm_provider": "anthropic"}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + repository.create.assert_awaited_once() + + def test_non_admin_cannot_create_a_credential_with_an_openai_token_file(self): + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ): + response = _post_credential( + { + "credential_name": "attacker-cred", + "credential_values": {"openai_identity_token_file": "/var/run/secrets/tokens/attacker"}, + "credential_info": {"custom_llm_provider": "openai"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "openai_identity_token_file" in response.json()["error"]["message"] + + def test_proxy_admin_can_create_a_credential_with_the_openai_identity_trio(self, restore_credential_list): + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "openai-wif", + "credential_values": { + "openai_identity_provider_id": "idp_1", + "openai_service_account_id": "user-1", + "openai_identity_token_file": "/var/run/secrets/tokens/openai", + }, + "credential_info": {"custom_llm_provider": "openai"}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + repository.create.assert_awaited_once() + + def test_non_admin_cannot_update_a_credential_to_add_a_wif_destination(self): + stored = CredentialItem( + credential_name="existing", + credential_values={"api_key": "sk-old"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.credential_endpoints.endpoints.CredentialsRepository" + ) as repository, # test-quality-ok: the proxy wiring under test is what this patches + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + update_mock = AsyncMock(return_value=None) + repository.return_value.update_by_name = update_mock + + response = _patch_credential( + "existing", + { + "credential_name": "existing", + "credential_values": {"anthropic_keycloak_token_url": "https://evil.example.com/token"}, + "credential_info": {}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + update_mock.assert_not_awaited() + + def test_non_admin_cannot_patch_wif_fields_onto_a_credential_through_model_id(self, credential_store): + """Regression: the PATCH gate read only the submitted ``credential_values``, so a non-admin + naming a federated deployment through ``model_id`` had its WIF fields copied onto an + ordinary credential unchecked, while POST already gated the resolved values.""" + stored = CredentialItem(credential_name="existing", credential_values={"api_key": "sk-old"}, credential_info={}) + update_by_name = AsyncMock(return_value=None) + router = MagicMock() + router.get_deployment.return_value = {"model_name": "claude-opus-5-5"} + router.get_deployment_credentials.return_value = { + "anthropic_keycloak_token_url": "https://keycloak.internal/token", + "anthropic_keycloak_client_secret_ref": "os.environ/KEYCLOAK_CLIENT_SECRET", + } + credential_store(find_by_name=AsyncMock(return_value=stored), update_by_name=update_by_name, llm_router=router) + + response = _patch_credential( + "existing", + {"credential_name": "existing", "model_id": "federated-deployment", "credential_info": {}}, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + update_by_name.assert_not_awaited() + + def test_proxy_admin_can_patch_wif_fields_onto_a_credential_through_model_id(self, credential_store): + stored = CredentialItem(credential_name="existing", credential_values={"api_key": "sk-old"}, credential_info={}) + update_by_name = AsyncMock(return_value=None) + router = MagicMock() + router.get_deployment.return_value = {"model_name": "claude-opus-5-5"} + router.get_deployment_credentials.return_value = {"anthropic_keycloak_token_url": "https://keycloak.internal/token"} + credential_store(find_by_name=AsyncMock(return_value=stored), update_by_name=update_by_name, llm_router=router) + + response = _patch_credential( + "existing", + {"credential_name": "existing", "model_id": "federated-deployment", "credential_info": {}}, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + written = json.loads(update_by_name.await_args.kwargs["data"]["credential_values"]) + assert "anthropic_keycloak_token_url" in written, "the deployment's WIF field reaches the stored credential" + + def test_proxy_admin_can_update_a_credential_to_add_a_wif_destination(self): + stored = CredentialItem( + credential_name="existing", + credential_values={"api_key": "sk-old"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", MagicMock() + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.credential_endpoints.endpoints.CredentialsRepository" + ) as repository, # test-quality-ok: the proxy wiring under test is what this patches + ): + repository.return_value.find_by_name = AsyncMock(return_value=stored) + update_mock = AsyncMock(return_value=None) + repository.return_value.update_by_name = update_mock + + response = _patch_credential( + "existing", + { + "credential_name": "existing", + "credential_values": {"anthropic_keycloak_token_url": "https://keycloak.internal/token"}, + "credential_info": {}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + update_mock.assert_awaited_once() + + +def _wif_credential(name: str = "federated-cred") -> CredentialItem: + return CredentialItem( + credential_name=name, + credential_values={ + "anthropic_keycloak_token_url": "https://keycloak.internal/token", + "api_key": "sk-old", + }, + credential_info={"custom_llm_provider": "anthropic"}, + ) + + +def _plain_credential(name: str = "ordinary-cred") -> CredentialItem: + return CredentialItem( + credential_name=name, + credential_values={"api_key": "sk-old"}, + credential_info={"custom_llm_provider": "openai"}, + ) + + +class TestNonAdminCannotTouchAStoredWifCredential: + """The WIF gate used to read only the incoming ``credential_values``, so a non-admin could + drop a federation field by naming it in ``credential_values_to_delete`` (breaking every + deployment that references the credential), or edit a stored admin-owned WIF credential by + sending a payload carrying no WIF field at all. The gate is evaluated against the effective + surface of the operation: incoming keys (a ``null`` value still persists the key), deleted + keys, and the stored credential, wherever it lives (DB row or config-only ``credential_list`` + entry).""" + + def test_non_admin_cannot_delete_a_wif_field_off_a_credential(self, restore_credential_list): + with _repository_holding(_plain_credential("some-cred")) as repository: + response = _patch_credential( + "some-cred", + { + "credential_name": "some-cred", + "credential_values": {}, + "credential_values_to_delete": ["anthropic_keycloak_token_url"], + "credential_info": {}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "anthropic_keycloak_token_url" in response.text + repository.update_by_name.assert_not_awaited() + + def test_non_admin_cannot_patch_a_stored_wif_credential(self, restore_credential_list): + with _repository_holding(_wif_credential("federated-cred")) as repository: + response = _patch_credential( + "federated-cred", + { + "credential_name": "federated-cred", + "credential_values": {"api_key": "sk-attacker"}, + "credential_info": {}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "anthropic_keycloak_token_url" in response.text + repository.update_by_name.assert_not_awaited() + + def test_proxy_admin_can_delete_a_wif_field_off_a_credential(self, restore_credential_list): + with _repository_holding(_wif_credential("federated-cred")) as repository: + response = _patch_credential( + "federated-cred", + { + "credential_name": "federated-cred", + "credential_values": {}, + "credential_values_to_delete": ["anthropic_keycloak_token_url"], + "credential_info": {}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + written_values = json.loads(repository.update_by_name.await_args.kwargs["data"]["credential_values"]) + assert "anthropic_keycloak_token_url" not in written_values + + def test_proxy_admin_can_patch_a_stored_wif_credential(self, restore_credential_list): + with _repository_holding(_wif_credential("federated-cred")) as repository: + response = _patch_credential( + "federated-cred", + { + "credential_name": "federated-cred", + "credential_values": {"api_key": "sk-rotated"}, + "credential_info": {}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + written_values = json.loads(repository.update_by_name.await_args.kwargs["data"]["credential_values"]) + assert written_values["anthropic_keycloak_token_url"] is not None + + def test_non_admin_can_still_patch_a_credential_with_no_wif_fields_anywhere(self, restore_credential_list): + with _repository_holding(_plain_credential("ordinary-cred")) as repository: + response = _patch_credential( + "ordinary-cred", + { + "credential_name": "ordinary-cred", + "credential_values": {"api_key": "sk-rotated"}, + "credential_info": {}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 200, response.text + repository.update_by_name.assert_awaited_once() + + def test_non_admin_cannot_delete_a_stored_wif_credential(self, restore_credential_list): + """DELETE takes the whole row, so it drops the admin-owned federation settings as surely + as a targeted key deletion would.""" + with _repository_holding(_wif_credential("federated-cred")) as repository: + response = _delete_credential("federated-cred", auth=_as_non_admin) + + assert response.status_code == 403, response.text + assert response.json()["error"]["param"] == "anthropic_keycloak_token_url" + repository.delete_by_name.assert_not_awaited() + + def test_a_stale_in_memory_copy_does_not_authorize_deleting_a_stored_wif_credential( + self, restore_credential_list, monkeypatch + ): + """Resolution reads memory first and stops, which is right when serving a request. A pod + whose in-memory copy predates an admin adding the federation fields must not read that + stale object and authorize the delete: the gate takes the union of memory and the row.""" + monkeypatch.setattr(litellm, "credential_list", [_plain_credential("federated-cred")]) + + with _repository_holding(_wif_credential("federated-cred")) as repository: + response = _delete_credential("federated-cred", auth=_as_non_admin) + + assert response.status_code == 403, response.text + assert response.json()["error"]["param"] == "anthropic_keycloak_token_url" + repository.delete_by_name.assert_not_awaited() + + def test_proxy_admin_can_delete_a_stored_wif_credential(self, restore_credential_list): + with _repository_holding(_wif_credential("federated-cred")) as repository: + response = _delete_credential("federated-cred", auth=_as_admin) + + assert response.status_code == 200, response.text + repository.delete_by_name.assert_awaited_once_with("federated-cred") + + def test_non_admin_can_still_delete_a_credential_with_no_wif_fields(self, restore_credential_list): + with _repository_holding(_plain_credential("ordinary-cred")) as repository: + response = _delete_credential("ordinary-cred", auth=_as_non_admin) + + assert response.status_code == 200, response.text + repository.delete_by_name.assert_awaited_once_with("ordinary-cred") + + def test_non_admin_cannot_null_out_a_wif_field_on_a_credential(self, restore_credential_list): + """A JSON ``null`` still lands as a key in ``credential_values``. ``get_litellm_params`` + forwards a WIF kwarg on key presence and the federation resolver rejects a foreign + variant's field by key, so a value-based gate let a non-admin persist the key and wedge + every deployment referencing the credential at request time.""" + with _repository_holding(_plain_credential("some-cred")) as repository: + response = _patch_credential( + "some-cred", + { + "credential_name": "some-cred", + "credential_values": {"anthropic_issuer_url": None}, + "credential_info": {}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "anthropic_issuer_url" in response.text + repository.update_by_name.assert_not_awaited() + + def test_non_admin_cannot_patch_a_credential_storing_a_null_wif_field(self, restore_credential_list): + stored = CredentialItem( + credential_name="nulled-cred", + credential_values={"anthropic_issuer_url": None, "api_key": "sk-old"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + with _repository_holding(stored) as repository: + response = _patch_credential( + "nulled-cred", + { + "credential_name": "nulled-cred", + "credential_values": {"api_key": "sk-attacker"}, + "credential_info": {}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "anthropic_issuer_url" in response.text + repository.update_by_name.assert_not_awaited() + + def test_proxy_admin_can_null_out_a_wif_field_on_a_credential(self, restore_credential_list): + with _repository_holding(_wif_credential("federated-cred")) as repository: + response = _patch_credential( + "federated-cred", + { + "credential_name": "federated-cred", + "credential_values": {"anthropic_keycloak_token_url": None}, + "credential_info": {}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + repository.update_by_name.assert_awaited_once() + + def test_non_admin_cannot_delete_a_config_only_wif_credential(self, restore_credential_list, monkeypatch): + """A ``credential_list`` entry from config.yaml has no DB row, so a gate that consulted + only the DB let a non-admin evict the admin-owned federation settings from memory.""" + config_credential = _wif_credential("config-wif") + monkeypatch.setattr(litellm, "credential_list", [config_credential]) + with _repository_holding(None) as repository: + response = _delete_credential("config-wif", auth=_as_non_admin) + + assert response.status_code == 403, response.text + assert response.json()["error"]["param"] == "anthropic_keycloak_token_url" + repository.delete_by_name.assert_not_awaited() + assert litellm.credential_list == [config_credential] + + def test_proxy_admin_can_delete_a_config_only_wif_credential(self, restore_credential_list, monkeypatch): + """The gate lets the admin through to the row delete. The 404 that follows is the rule for + every config-only credential (no row to delete, the entry is back on the next boot), so the + in-memory entry stays put too.""" + config_credential = _wif_credential("config-wif") + monkeypatch.setattr(litellm, "credential_list", [config_credential]) + with _repository_holding(None) as repository: + response = _delete_credential("config-wif", auth=_as_admin) + + assert response.status_code == 404, response.text + repository.delete_by_name.assert_awaited_once_with("config-wif") + assert litellm.credential_list == [config_credential] + + def test_non_admin_cannot_shadow_a_config_only_wif_credential(self, restore_credential_list, monkeypatch): + """POST with the same name carries no WIF field and collides with no DB row, yet + ``CredentialAccessor.upsert_credentials`` would replace the admin entry in memory and + the periodic config sync would then make the takeover permanent.""" + config_credential = _wif_credential("config-wif") + monkeypatch.setattr(litellm, "credential_list", [config_credential]) + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "config-wif", + "credential_values": {"api_key": "sk-attacker"}, + "credential_info": {"custom_llm_provider": "anthropic"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "anthropic_keycloak_token_url" in response.text + repository.create.assert_not_awaited() + assert litellm.credential_list == [config_credential] + assert litellm.credential_list[0].credential_values["api_key"] == "sk-old" + + def test_proxy_admin_can_post_over_a_config_only_wif_credential(self, restore_credential_list, monkeypatch): + monkeypatch.setattr(litellm, "credential_list", [_wif_credential("config-wif")]) + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "config-wif", + "credential_values": {"api_key": "sk-rotated"}, + "credential_info": {"custom_llm_provider": "anthropic"}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + repository.create.assert_awaited_once() + assert litellm.credential_list[0].credential_values == {"api_key": "sk-rotated"} + + def test_non_admin_cannot_rename_a_credential_onto_a_config_only_wif_credential( + self, restore_credential_list, monkeypatch + ): + """PATCH is the other way to shadow: renaming an ordinary credential onto the WIF + credential's name makes ``_sync_in_memory_credential`` upsert the attacker's values over + the admin entry, with no WIF field in the payload and no DB row to collide with.""" + config_credential = _wif_credential("config-wif") + monkeypatch.setattr(litellm, "credential_list", [_plain_credential("mine"), config_credential]) + with _repository_holding(_plain_credential("mine")) as repository: + response = _patch_credential( + "mine", + { + "credential_name": "config-wif", + "credential_values": {"api_key": "sk-attacker"}, + "credential_info": {}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "anthropic_keycloak_token_url" in response.text + repository.update_by_name.assert_not_awaited() + assert config_credential in litellm.credential_list + assert litellm.credential_list[1].credential_values["api_key"] == "sk-old" + + def test_proxy_admin_can_rename_a_credential_onto_a_config_only_wif_credential( + self, restore_credential_list, monkeypatch + ): + monkeypatch.setattr(litellm, "credential_list", [_plain_credential("mine"), _wif_credential("config-wif")]) + with _repository_holding(_plain_credential("mine")) as repository: + response = _patch_credential( + "mine", + { + "credential_name": "config-wif", + "credential_values": {"api_key": "sk-rotated"}, + "credential_info": {}, + }, + auth=_as_admin, + ) + + assert response.status_code == 200, response.text + repository.update_by_name.assert_awaited_once() + assert [c.credential_name for c in litellm.credential_list] == ["config-wif"] + + def test_non_admin_cannot_post_a_null_wif_field(self, restore_credential_list): + """Same key-presence rule on the create path: ``{"anthropic_issuer_url": null}`` persists + the key, and the resolver reacts to the key.""" + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "nulled-cred", + "credential_values": {"anthropic_issuer_url": None, "api_key": "sk-new"}, + "credential_info": {"custom_llm_provider": "anthropic"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + assert "anthropic_issuer_url" in response.text + repository.create.assert_not_awaited() + assert litellm.credential_list == [] + + def test_non_admin_cannot_shadow_a_db_stored_wif_credential(self, restore_credential_list): + """Same hole for a WIF credential another pod wrote to the DB before this pod's in-memory + list caught up: the existing-credential lookup falls through to the DB.""" + with _repository_holding(_wif_credential("federated-cred")) as repository: + response = _post_credential( + { + "credential_name": "federated-cred", + "credential_values": {"api_key": "sk-attacker"}, + "credential_info": {"custom_llm_provider": "anthropic"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 403, response.text + repository.create.assert_not_awaited() + + def test_non_admin_can_still_post_a_credential_with_no_wif_fields_anywhere(self, restore_credential_list): + with _repository_holding(None) as repository: + response = _post_credential( + { + "credential_name": "ordinary-cred", + "credential_values": {"api_key": "sk-new"}, + "credential_info": {"custom_llm_provider": "openai"}, + }, + auth=_as_non_admin, + ) + + assert response.status_code == 200, response.text + repository.create.assert_awaited_once() + assert litellm.credential_list[0].credential_name == "ordinary-cred" + + +class TestManagementReadsTheStoredCredential: + """Serving a request reads memory first, which is right. A management operation cannot: on a + pod whose in-memory copy predates another pod's update it would act on superseded values.""" + + @pytest.mark.asyncio + async def test_authoritative_hydrate_prefers_the_row_over_a_stale_memory_copy(self): + import litellm + from litellm.proxy.common_utils.credential_hydration import ( + hydrate_named_credential, + hydrate_named_credential_authoritative, + ) + from litellm.types.utils import CredentialItem + + stale = CredentialItem( + credential_name="anthropic-wif", + credential_values={"anthropic_issuer_url": "https://old.example.com"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + row = { + "credential_name": "anthropic-wif", + "credential_values": {"anthropic_issuer_url": "https://new.example.com"}, + "credential_info": {"custom_llm_provider": "anthropic"}, + } + + prisma = MagicMock() + prisma.db.litellm_credentialstable.find_unique = AsyncMock(return_value=row) + + with patch.object(litellm, "credential_list", [stale]): # test-quality-ok: the stale copy under test + served = await hydrate_named_credential("anthropic-wif", prisma) + managed = await hydrate_named_credential_authoritative("anthropic-wif", prisma) + + assert served is not None and served.credential_values["anthropic_issuer_url"] == "https://old.example.com" + assert managed is not None and managed.credential_values["anthropic_issuer_url"] == "https://new.example.com" + + @pytest.mark.asyncio + async def test_authoritative_hydrate_falls_back_to_memory_when_the_row_is_absent(self): + import litellm + from litellm.proxy.common_utils.credential_hydration import hydrate_named_credential_authoritative + from litellm.types.utils import CredentialItem + + only_in_memory = CredentialItem( + credential_name="config-yaml-credential", + credential_values={"anthropic_issuer_url": "https://configured.example.com"}, + credential_info={"custom_llm_provider": "anthropic"}, + ) + prisma = MagicMock() + prisma.db.litellm_credentialstable.find_unique = AsyncMock(return_value=None) + + with patch.object(litellm, "credential_list", [only_in_memory]): # test-quality-ok: the config.yaml fallback under test + resolved = await hydrate_named_credential_authoritative("config-yaml-credential", prisma) + + assert resolved is not None + assert resolved.credential_values["anthropic_issuer_url"] == "https://configured.example.com" + + def test_delete_credential_answers_404_when_the_credential_does_not_exist(credential_store): """Regression: prisma's ``delete`` hands back None when the ``where`` clause matched no row instead of raising, and the handler never looked. Deleting a name that was never stored @@ -131,7 +1200,7 @@ def test_delete_credential_answers_404_when_the_credential_does_not_exist(creden assert response.status_code == 404, ( f"delete of a missing credential answered {response.status_code}: {response.text}" ) - assert "definitely-not-there" in response.text + assert "definitely-not-there" in response.json()["error"]["message"] def test_delete_credential_still_answers_200_and_drops_the_credential_from_memory(credential_store): @@ -206,7 +1275,7 @@ def test_get_credentials_answers_an_error_status_when_the_listing_fails(credenti def _create_credential(body: dict): - return _call_as_admin("POST", "/credentials", body) + return _call_as("POST", "/credentials", body) class _UniqueViolation(Exception): diff --git a/tests/unit/proxy/health_endpoints/test_health_endpoints.py b/tests/unit/proxy/health_endpoints/test_health_endpoints.py index 513f4a67554..95aad7b772f 100644 --- a/tests/unit/proxy/health_endpoints/test_health_endpoints.py +++ b/tests/unit/proxy/health_endpoints/test_health_endpoints.py @@ -33,6 +33,11 @@ from litellm.proxy.health_endpoints._health_endpoints import ( from litellm.proxy.health_endpoints._health_endpoints import ( test_model_connection as health_test_model_connection, ) +from litellm.types.workload_identity import ( + ANTHROPIC_WIF_KWARGS_KEYS, + OPENAI_WIF_KWARGS_KEYS, + WIF_SECRET_BEARING_KEYS, +) # Import shared proxy test helpers from conftest from tests.unit.proxy.conftest import create_proxy_test_client @@ -1067,6 +1072,112 @@ async def test_test_model_connection_uses_loaded_deployment_team_id_via_model_na assert passed_model_params.model_info.team_id == deployment_owner_team_id +FEDERATED_DEPLOYMENT_ID = "federated-deployment-id" +FEDERATED_DEPLOYMENT_TEAM_ID = "team-owning-the-federated-deployment" + + +def _federated_deployment(): + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo + + return Deployment( + model_name="claude-federated", + litellm_params=LiteLLM_Params( + model="anthropic/claude-sonnet-4-5", + api_base="https://api.anthropic.com", + anthropic_federation_rule_id="rule-abc", + anthropic_organization_id="org-abc", + anthropic_identity_source="oidc/env/PROXY_OIDC_TOKEN", + ), + model_info=ModelInfo(id=FEDERATED_DEPLOYMENT_ID, team_id=FEDERATED_DEPLOYMENT_TEAM_ID), + ) + + +async def _probe_federated_deployment_as_team_admin(litellm_params): + """Run the Test Connection button against a federated deployment as an admin of its own team. + + ``allow_client_side_credentials`` is on, which is what lets a request-supplied api_base keep + the configuration's credentials instead of dropping them, so the federation params are still + on the deployment being probed when the request redirects it. + """ + from litellm.proxy._types import LiteLLM_TeamTable + + mock_router = MagicMock() + mock_router.get_deployment.return_value = _federated_deployment() + + async def fake_find_unique(*, where): + if where["team_id"] != FEDERATED_DEPLOYMENT_TEAM_ID: + return None + return SimpleNamespace( + model_dump=lambda: LiteLLM_TeamTable( + team_id=FEDERATED_DEPLOYMENT_TEAM_ID, + members_with_roles=[{"user_id": "team-admin-user", "role": "admin"}], + ).model_dump() + ) + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_teamtable.find_unique = AsyncMock(side_effect=fake_find_unique) + mock_ahealth_check = AsyncMock(return_value={"status": "healthy"}) + + with ( + patch.multiple( # test-quality-ok: proxy module globals, no injection seam + "litellm.proxy.proxy_server", + prisma_client=mock_prisma_client, + llm_router=mock_router, + premium_user=True, + general_settings={"allow_client_side_credentials": True}, + ), + patch( # test-quality-ok: the probe params handed to the health check are the assertion + "litellm.proxy.health_endpoints._health_endpoints.litellm.ahealth_check", + mock_ahealth_check, + ), + ): + response = await health_test_model_connection( + request=MagicMock(), + mode="chat", + litellm_params=litellm_params, + model_info={"id": FEDERATED_DEPLOYMENT_ID}, + user_api_key_dict=UserAPIKeyAuth( + token="requester-token", + user_id="team-admin-user", + team_id=FEDERATED_DEPLOYMENT_TEAM_ID, + user_role=LitellmUserRoles.INTERNAL_USER, + ), + ) + return response, mock_ahealth_check + + +@pytest.mark.asyncio +async def test_test_connection_still_lets_a_team_admin_probe_a_federated_deployment(): + """A probe that changes nothing about the deployment is not a credential change, so the team + admin who owns the deployment can still press Test Connection on it.""" + response, mock_ahealth_check = await _probe_federated_deployment_as_team_admin( + {"model": "anthropic/claude-sonnet-4-5"} + ) + + assert response["status"] == "success" + assert mock_ahealth_check.await_count == 1 + assert mock_ahealth_check.await_args.kwargs["model_params"]["anthropic_federation_rule_id"] == "rule-abc" + + +@pytest.mark.asyncio +async def test_test_connection_refuses_a_non_admin_pointing_a_federated_deployment_elsewhere(): + """A probe carrying its own api_base sends the deployment's minted org-scoped token to a host + the caller chose, so it is a credential change and only a proxy admin may make it. The probe + used to authorize with nothing declared as incoming, which left this gate unreachable here.""" + from litellm.proxy._types import ProxyException + + with pytest.raises(ProxyException) as exc_info: + await _probe_federated_deployment_as_team_admin( + { + "model": "anthropic/claude-sonnet-4-5", + "api_base": "https://caller-chosen.invalid/v1", + } + ) + + assert exc_info.value.code == "403" + assert "workload identity federation" in exc_info.value.message + + @pytest.mark.asyncio async def test_test_model_connection_authorizes_on_params_after_health_check_params_merge(): """ @@ -2358,10 +2469,11 @@ async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not(): # Non-admin response must advertise that api_base/api_version were # withheld so clients that previously parsed them can detect the change. - assert ( - non_admin_response.headers.get("Litellm-Health-Field-Notice") - == "api_base, api_version, aws_bedrock_runtime_endpoint are admin-only on this endpoint" - ) + notice = non_admin_response.headers.get("Litellm-Health-Field-Notice") + assert notice is not None + withheld = notice.removesuffix(" are admin-only on this endpoint").split(", ") + assert {"api_base", "api_version", "aws_bedrock_runtime_endpoint"} <= set(withheld) + assert [field for field in withheld if field in non_admin_eps[0]] == [] assert "Litellm-Health-Field-Notice" not in admin_response.headers # Stripping must produce a copy — the shared cache must still carry the @@ -2371,6 +2483,127 @@ async def test_health_endpoint_admin_sees_routing_fields_non_admin_does_not(): assert cached_first["api_version"] == "2024-10-21" +@pytest.mark.asyncio +async def test_health_endpoint_keeps_federation_identity_admin_only(): + """A federated deployment's health entry names the identity it mints as: the rule, the workspace, + the service account, the issuer it signs against. That is the same routing detail api_base is, + so a non-admin who can see the deployment is healthy must not learn which identity it borrows, + and an admin debugging a failing exchange must still see all of it. + """ + from fastapi import Response + + from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth + from litellm.proxy.health_endpoints._health_endpoints import health_endpoint + + federation_fields = { + "anthropic_federation_rule_id": "fdrl_01H", + "anthropic_federation_workspace_id": "wrkspc_01H", + "anthropic_organization_id": "org-acme", + "anthropic_service_account_id": "svc_01H", + "anthropic_identity_source": "oidc/env/OIDC_TOKEN", + "anthropic_issuer_url": "https://issuer.internal", + "anthropic_keycloak_client_id": "litellm-proxy", + "openai_identity_provider_id": "idp_01H", + "openai_service_account_id": "sa_01H", + } + full_model_list = [ + { + "model_name": "model-a", + "litellm_params": {"model": "anthropic/claude-sonnet-5", **federation_fields}, + "model_info": {"id": "id-a"}, + }, + ] + cached_results = { + "healthy_endpoints": [{"model": "anthropic/claude-sonnet-5", "model_id": "id-a", **federation_fields}], + "unhealthy_endpoints": [], + "healthy_count": 1, + "unhealthy_count": 0, + } + + admin_key = UserAPIKeyAuth( + api_key="hashed-admin-key", + models=["model-a"], + user_role=LitellmUserRoles.PROXY_ADMIN, + ) + non_admin_key = UserAPIKeyAuth(api_key="hashed-user-key", models=["model-a"]) + + with patch.multiple( # test-quality-ok: proxy module globals, no injection seam + "litellm.proxy.proxy_server", + llm_model_list=full_model_list, + llm_router=None, + prisma_client=None, + use_background_health_checks=True, + user_model=None, + health_check_results=cached_results, + health_check_details=True, + health_check_concurrency=1, + ): + admin_result = await health_endpoint( + response=Response(), + user_api_key_dict=admin_key, + model=None, + model_id=None, + ) + non_admin_result = await health_endpoint( + response=Response(), + user_api_key_dict=non_admin_key, + model=None, + model_id=None, + ) + + admin_endpoint = admin_result["healthy_endpoints"][0] + non_admin_endpoint = non_admin_result["healthy_endpoints"][0] + + assert {key: admin_endpoint.get(key) for key in federation_fields} == federation_fields + assert [key for key in federation_fields if key in non_admin_endpoint] == [] + assert non_admin_endpoint["model_id"] == "id-a" + + +@pytest.mark.parametrize("federation_field", sorted(ANTHROPIC_WIF_KWARGS_KEYS | OPENAI_WIF_KWARGS_KEYS)) +def test_no_federation_field_reaches_a_non_admin_health_entry(federation_field: str): + """Every key that configures workload identity federation either names the identity a + deployment mints as or carries the secret it mints with, and a non-admin who can see the + deployment is healthy must learn neither. Both lists that enforce that are derived from the + same key sets this runs over, so a field added to the funnel without joining either one shows + up here as a value a non-admin could read.""" + from litellm.proxy.health_check import _clean_endpoint_data + from litellm.proxy.health_endpoints._health_endpoints import ( + _strip_admin_only_fields_from_health_result, + ) + + canary = f"CANARY-{federation_field}-VALUE" + cleaned = _clean_endpoint_data( + {"model": "anthropic/claude-sonnet-5", federation_field: canary}, + details=True, + ) + stripped = _strip_admin_only_fields_from_health_result( + {"healthy_endpoints": [cleaned], "unhealthy_endpoints": []} + ) + + assert stripped["healthy_endpoints"][0]["model"] == "anthropic/claude-sonnet-5" + assert federation_field not in stripped["healthy_endpoints"][0] + assert canary not in str(stripped) + + +@pytest.mark.parametrize("secret_field", sorted(WIF_SECRET_BEARING_KEYS)) +def test_no_federation_secret_reaches_even_an_admin_health_entry(secret_field: str): + """A proxy admin is allowed to read which identity a deployment federates as, but never the + token, key, or reference it federates with, so these fields drop at the health-check layer + ahead of any per-caller stripping. Reading the same set the drop list is built from is what + catches a new secret-bearing field that was only ever added to the admin-gated half.""" + from litellm.proxy.health_check import _clean_endpoint_data + + canary = f"CANARY-{secret_field}-VALUE" + cleaned = _clean_endpoint_data( + {"model": "anthropic/claude-sonnet-5", secret_field: canary}, + details=True, + ) + + assert cleaned["model"] == "anthropic/claude-sonnet-5" + assert secret_field not in cleaned + assert canary not in str(cleaned) + + @pytest.mark.asyncio async def test_health_endpoint_warns_when_scoped_models_lack_model_id(): """ @@ -3189,6 +3422,11 @@ def test_clean_endpoint_data_strips_extra_headers_and_aws_session_token(): "aws_secret_access_key", "aws_session_token", "aws_web_identity_token", + "anthropic_identity_token", + "anthropic_issuer_signing_key_ref", + "anthropic_keycloak_client_secret_ref", + "anthropic_identity_token_file", + "openai_identity_token_file", "vertex_credentials", "vertex_ai_credentials", "extra_headers", diff --git a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py index 19cc8bf15ab..8159890ef16 100644 --- a/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_model_management_endpoints.py @@ -65,11 +65,7 @@ class MockPrismaClient: return LiteLLM_TeamTable( team_id=where["team_id"], team_alias="test_team", - members_with_roles=[ - Member( - user_id="test_user", role="admin" if self.user_admin else "user" - ) - ], + members_with_roles=[Member(user_id="test_user", role="admin" if self.user_admin else "user")], ) return None @@ -83,10 +79,7 @@ class MockPrismaClient: # Support model_name startswith filter (used by _get_team_deployments) if where and "model_name" in where: model_name_filter = where["model_name"] - if ( - isinstance(model_name_filter, dict) - and "startswith" in model_name_filter - ): + if isinstance(model_name_filter, dict) and "startswith" in model_name_filter: prefix = model_name_filter["startswith"] results = [d for d in results if d.model_name.startswith(prefix)] @@ -131,13 +124,9 @@ class MockProxyConfig: class TestModelManagementAuthChecks: def setup_method(self): """Setup test cases""" - self.admin_user = UserAPIKeyAuth( - user_id="test_admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + self.admin_user = UserAPIKeyAuth(user_id="test_admin", user_role=LitellmUserRoles.PROXY_ADMIN) - self.normal_user = UserAPIKeyAuth( - user_id="test_user", user_role=LitellmUserRoles.INTERNAL_USER - ) + self.normal_user = UserAPIKeyAuth(user_id="test_user", user_role=LitellmUserRoles.INTERNAL_USER) self.team_admin_user = UserAPIKeyAuth( user_id="test_user", @@ -156,7 +145,7 @@ class TestModelManagementAuthChecks: @pytest.mark.asyncio async def test_can_user_make_team_model_call_non_premium_fails(self): """Test that non-premium users cannot make team model calls""" - with pytest.raises(Exception, match='You must be a LiteLLM Enterprise user to use this feature\\.') as exc_info: + with pytest.raises(Exception, match="You must be a LiteLLM Enterprise user to use this feature\\.") as exc_info: ModelManagementAuthChecks.can_user_make_team_model_call( team_id="test_team", user_api_key_dict=self.admin_user, @@ -170,9 +159,7 @@ class TestModelManagementAuthChecks: team_obj = LiteLLM_TeamTable( team_id="test_team", team_alias="test_team", - members_with_roles=[ - Member(user_id=self.team_admin_user.user_id, role="admin") - ], + members_with_roles=[Member(user_id=self.team_admin_user.user_id, role="admin")], ) result = ModelManagementAuthChecks.can_user_make_team_model_call( @@ -211,7 +198,7 @@ class TestModelManagementAuthChecks: ) prisma_client = MockPrismaClient(team_exists=True) - with pytest.raises(Exception, match='You must be a LiteLLM Enterprise user to use this feature\\.') as exc_info: + with pytest.raises(Exception, match="You must be a LiteLLM Enterprise user to use this feature\\.") as exc_info: await ModelManagementAuthChecks.allow_team_model_action( model_params=model_params, user_api_key_dict=self.admin_user, @@ -258,6 +245,7 @@ class TestModelManagementAuthChecks: user_api_key_dict=self.admin_user, prisma_client=prisma_client, premium_user=True, + incoming_params=None, ) assert result is True @@ -279,6 +267,7 @@ class TestModelManagementAuthChecks: user_api_key_dict=self.normal_user, prisma_client=prisma_client, premium_user=True, + incoming_params=None, ) assert "403" in str(exc_info.value) @@ -706,29 +695,21 @@ class TestDeleteTeamModelAlias: mock_prisma.db = MockPrismaWrapper(model_aliases_list) # Call the function - await delete_team_model_alias( - public_model_name="public_model_1", prisma_client=mock_prisma - ) + await delete_team_model_alias(public_model_name="public_model_1", prisma_client=mock_prisma) # Verify results mock_db = mock_prisma.db.litellm_modeltable - assert ( - len(mock_db.update_calls) == 2 - ) # Should have 2 update calls since public_model_1 appears twice + assert len(mock_db.update_calls) == 2 # Should have 2 update calls since public_model_1 appears twice # Verify first update first_update = mock_db.update_calls[0] assert first_update["where"] == {"id": 1} - assert json.loads(first_update["data"]["model_aliases"]) == { - "alias2": "public_model_2" - } + assert json.loads(first_update["data"]["model_aliases"]) == {"alias2": "public_model_2"} # Verify second update second_update = mock_db.update_calls[1] assert second_update["where"] == {"id": 2} - assert json.loads(second_update["data"]["model_aliases"]) == { - "alias3": "public_model_3" - } + assert json.loads(second_update["data"]["model_aliases"]) == {"alias3": "public_model_3"} @pytest.mark.asyncio async def test_delete_team_model_alias_no_matches(self): @@ -764,9 +745,7 @@ class TestDeleteTeamModelAlias: mock_prisma.db = MockPrismaWrapper(model_aliases_list) # Call the function with non-existent model - await delete_team_model_alias( - public_model_name="non_existent_model", prisma_client=mock_prisma - ) + await delete_team_model_alias(public_model_name="non_existent_model", prisma_client=mock_prisma) # Verify no updates were made mock_db = mock_prisma.db.litellm_modeltable @@ -1376,18 +1355,12 @@ class TestUpdateModel: updated_row.model_dump_json.return_value = "{}" mock_prisma = MagicMock() - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=existing_row - ) - mock_prisma.db.litellm_proxymodeltable.update = AsyncMock( - return_value=updated_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) mock_router = MagicMock() mock_router.get_model_ids.return_value = [model_id] - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) with ( patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), @@ -1404,9 +1377,7 @@ class TestUpdateModel: ), patch( # test-quality-ok: [TQ008] isolate persistence from router reload implementation "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", - new=AsyncMock( - return_value=ReconcileOutcome(still_desired=None, live_after=None) - ), + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), ) as mock_clear_cache, ): await update_model( @@ -1505,9 +1476,7 @@ class TestUpdatePublicModelGroups: mock_proxy_config.get_config = mock_get_config mock_proxy_config.save_config = AsyncMock() - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) request = UpdatePublicModelGroupsRequest(model_groups=new_models) @@ -1563,9 +1532,7 @@ class TestUpdatePublicModelGroups: mock_proxy_config.get_config = mock_get_config mock_proxy_config.save_config = AsyncMock() - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) request = UpdateUsefulLinksRequest(useful_links=new_links) @@ -1730,9 +1697,7 @@ class TestTeamModelSiblingRouting: ) # Global deployment should be accessible when team_id is provided - deployments = router._get_all_deployments( - model_name="global-gpt-4o", team_id="teamA" - ) + deployments = router._get_all_deployments(model_name="global-gpt-4o", team_id="teamA") assert len(deployments) == 1 assert deployments[0]["model_name"] == "global-gpt-4o" @@ -1781,9 +1746,9 @@ class TestTeamModelUpdate: patch( "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" ) as mock_team_model_add, - patch( + patch( # test-quality-ok: the proxy wiring under test is what this patches "litellm.proxy.management_endpoints.model_management_endpoints.update_team" - ) as mock_update_team, + ) as mock_update_team, # test-quality-ok: the proxy wiring under test is what this patches ): result = await _update_team_model_in_db( db_model=db_model, @@ -1814,9 +1779,7 @@ class TestTeamModelUpdate: db_model = Deployment( model_name="model_name_team_123_uuid1", litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo( - team_id="team_123", team_public_model_name="old-public-name" - ), + model_info=ModelInfo(team_id="team_123", team_public_model_name="old-public-name"), ) # Create a sibling deployment that still uses the old public name @@ -1827,9 +1790,7 @@ class TestTeamModelUpdate: "team_public_model_name": "old-public-name", } - prisma_client = MockPrismaClient( - team_exists=True, sibling_deployments=[sibling_deployment] - ) + prisma_client = MockPrismaClient(team_exists=True, sibling_deployments=[sibling_deployment]) patch_data = updateDeployment( model_name="new-public-name", @@ -1844,10 +1805,10 @@ class TestTeamModelUpdate: with ( patch( "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, + ) as mock_delete, # test-quality-ok: the proxy wiring under test is what this patches patch( "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add, + ) as mock_add, # test-quality-ok: the proxy wiring under test is what this patches ): await _update_existing_team_model_assignment( team_id="team_123", @@ -1889,10 +1850,10 @@ class TestTeamModelUpdate: with ( patch( "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, + ) as mock_delete, # test-quality-ok: the proxy wiring under test is what this patches patch( "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add, + ) as mock_add, # test-quality-ok: the proxy wiring under test is what this patches ): await _update_existing_team_model_assignment( team_id="team_123", @@ -1910,7 +1871,6 @@ class TestTeamModelUpdate: """The team's model list autocommits, so it is written only after the row write succeeded: a refused write (the heuristic_v2 slot 403, a DB error) must not leave the team listing a name whose row never changed.""" - from fastapi import HTTPException from litellm.proxy.management_endpoints.model_management_endpoints import ( _update_team_model_in_db, @@ -1986,20 +1946,14 @@ class TestTeamModelUpdate: db_model = Deployment( model_name="model_name_team_123_uuid1", litellm_params=LiteLLM_Params(model="azure/gpt-4o-mini"), - model_info=ModelInfo( - team_id="team_123", team_public_model_name="old-public-name" - ), + model_info=ModelInfo(team_id="team_123", team_public_model_name="old-public-name"), ) sibling_deployment = MagicMock() sibling_deployment.model_name = "model_name_team_123_uuid2" - sibling_deployment.model_info = ( - '{"team_id":"team_123","team_public_model_name":"old-public-name"}' - ) + sibling_deployment.model_info = '{"team_id":"team_123","team_public_model_name":"old-public-name"}' - prisma_client = MockPrismaClient( - team_exists=True, sibling_deployments=[sibling_deployment] - ) + prisma_client = MockPrismaClient(team_exists=True, sibling_deployments=[sibling_deployment]) patch_data = updateDeployment( model_name="new-public-name", @@ -2014,10 +1968,10 @@ class TestTeamModelUpdate: with ( patch( "litellm.proxy.management_endpoints.model_management_endpoints.team_model_delete" - ) as mock_delete, + ) as mock_delete, # test-quality-ok: the proxy wiring under test is what this patches patch( "litellm.proxy.management_endpoints.model_management_endpoints.team_model_add" - ) as mock_add, + ) as mock_add, # test-quality-ok: the proxy wiring under test is what this patches ): await _update_existing_team_model_assignment( team_id="team_123", @@ -2094,10 +2048,7 @@ class TestTeamModelUpdate: ), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" def test_get_public_model_name_preserves_db_public_name_when_internal_name_unchanged( self, @@ -2124,10 +2075,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" def test_get_public_model_name_allows_top_level_rename(self): """A genuine rename via the top-level model_name field (no @@ -2152,10 +2100,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "new-public-name" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "new-public-name" def test_get_public_model_name_top_level_rename_wins_over_stale_model_info(self): """Regression (codex review): on a dashboard rename the UI sends the new @@ -2172,9 +2117,7 @@ class TestTeamModelUpdate: db_model = Deployment( model_name="model_name_team-a_abc123", litellm_params=LiteLLM_Params(model="azure/gpt-4.1"), - model_info=ModelInfo( - team_id="team-a", team_public_model_name="old-public-name" - ), + model_info=ModelInfo(team_id="team-a", team_public_model_name="old-public-name"), ) patch_data = updateDeployment( model_name="new-public-name", @@ -2184,10 +2127,7 @@ class TestTeamModelUpdate: ), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "new-public-name" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "new-public-name" def test_get_public_model_name_falls_back_to_db_public_name(self): """When patch_data carries no name hints at all (neither model_name @@ -2210,10 +2150,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" def test_get_public_model_name_last_resort_returns_db_model_name(self): """Legacy rows may have no team_public_model_name anywhere; the @@ -2233,10 +2170,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "legacy-model" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "legacy-model" def test_get_public_model_name_ignores_different_internal_shape_name(self): """A stale client may PATCH an internal-shaped model_name that does not @@ -2260,10 +2194,7 @@ class TestTeamModelUpdate: model_info=ModelInfo(team_id="test-team"), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" def test_get_public_model_name_ignores_internal_shape_patch_public(self): """If a corrupted row round-trips an internal-shaped value in @@ -2289,10 +2220,7 @@ class TestTeamModelUpdate: ), ) - assert ( - _get_public_model_name(patch_data=patch_data, db_model=db_model) - == "gpt-5.2-low-rpm-testing" - ) + assert _get_public_model_name(patch_data=patch_data, db_model=db_model) == "gpt-5.2-low-rpm-testing" @pytest.mark.asyncio async def test_dashboard_edit_preserves_public_name_and_acl(self): @@ -2360,9 +2288,7 @@ class TestTeamModelUpdate: # the merged model_info written to the DB must keep the public name model_info_json = result.get("model_info", "") parsed_model_info = json.loads(model_info_json) - assert ( - parsed_model_info.get("team_public_model_name") == "gpt-5.2-low-rpm-testing" - ) + assert parsed_model_info.get("team_public_model_name") == "gpt-5.2-low-rpm-testing" # the internal model_name must not have been overwritten (caller # intentionally clears patch_data.model_name so the DB row's name @@ -2404,9 +2330,7 @@ class TestModelInfoEndpoint: model_info=ModelInfo(id="gpt-4"), ) - result = await model_info( - model_id="gpt-4", user_api_key_dict=user_api_key_dict - ) + result = await model_info(model_id="gpt-4", user_api_key_dict=user_api_key_dict) assert result["id"] == "gpt-4" assert result["object"] == "model" @@ -2416,7 +2340,6 @@ class TestModelInfoEndpoint: @pytest.mark.asyncio async def test_model_info_inaccessible_model_returns_404(self): """Test model_info returns 404 for inaccessible models""" - from fastapi import HTTPException from litellm.proxy.proxy_server import model_info @@ -2481,9 +2404,7 @@ class TestModelInfoEndpoint: model_info=ModelInfo(id="team-model-1"), ) - result = await model_info( - model_id="team-model-1", user_api_key_dict=user_api_key_dict - ) + result = await model_info(model_id="team-model-1", user_api_key_dict=user_api_key_dict) assert result["id"] == "team-model-1" assert result["object"] == "model" @@ -2515,9 +2436,7 @@ class TestAddAndDeleteModelLifecycle: ) model_id = "lifecycle-test-model-123" - admin_user = UserAPIKeyAuth( - user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="test-admin", user_role=LitellmUserRoles.PROXY_ADMIN) # Build a real LiteLLM_ProxyModelTable for the DB mock to return db_row = LiteLLM_ProxyModelTable( @@ -2534,9 +2453,7 @@ class TestAddAndDeleteModelLifecycle: mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=db_row) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_proxy_config = MagicMock() @@ -2560,14 +2477,11 @@ class TestAddAndDeleteModelLifecycle: patch(f"{_PS}.llm_router", mock_router), patch(_ENCRYPT, side_effect=lambda value, **kwargs: value), ): - # --- ADD --- add_result = await add_new_model( model_params=Deployment( model_name="lifecycle-model", - litellm_params=LiteLLM_Params( - model="openai/gpt-4.1-nano", api_key="fake-key" - ), + litellm_params=LiteLLM_Params(model="openai/gpt-4.1-nano", api_key="fake-key"), model_info={"id": model_id}, ), user_api_key_dict=admin_user, @@ -2582,9 +2496,7 @@ class TestAddAndDeleteModelLifecycle: assert "deleted successfully" in delete_result["message"] # --- DELETE again should fail (model not found) --- - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=None - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=None) from litellm.proxy.proxy_server import ProxyException with pytest.raises(ProxyException) as exc_info: @@ -2646,24 +2558,18 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) # After the row delete no team deployment remains -> nothing backs the public name. mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) - mock_prisma.db.litellm_teamtable.update = AsyncMock( - return_value=updated_team_row - ) + mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=updated_team_row) # Team BYOK models have no alias row; delete_team_model_alias finds nothing. mock_prisma.db.litellm_modeltable = AsyncMock() mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2729,9 +2635,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() @@ -2741,9 +2645,7 @@ class TestDeleteTeamBYOKModelGhost: # No alias row matches -> delete_team_model_alias returns nothing, but it still ran. mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2806,25 +2708,17 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=deleted_row - ) - mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock( - return_value=deleted_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=deleted_row) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=deleted_row) # After the deleted replica's row is gone, the sibling still backs the public name. - mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock( - return_value=[sibling_row] - ) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[sibling_row]) mock_prisma.db.litellm_teamtable = AsyncMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=team_row) mock_prisma.db.litellm_modeltable = AsyncMock() mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2882,9 +2776,7 @@ class TestDeleteTeamBYOKModelGhost: members_with_roles=[Member(user_id="admin", role="admin")], models=[public_name], ) - alias_row = MagicMock( - id="alias-row-1", model_aliases={public_name: internal_name} - ) + alias_row = MagicMock(id="alias-row-1", model_aliases={public_name: internal_name}) alias_row.team = MagicMock() alias_row.team.team_id = team_id @@ -2892,26 +2784,20 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) mock_prisma.db.litellm_teamtable.update = AsyncMock(return_value=team_row) mock_prisma.db.litellm_modeltable = AsyncMock() - mock_prisma.db.litellm_modeltable.find_many = AsyncMock( - return_value=[alias_row] - ) + mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[alias_row]) mock_prisma.db.litellm_modeltable.update = AsyncMock() mock_router = MagicMock() mock_router.model_name_to_deployment_indices = {public_name: [0]} - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -2974,9 +2860,7 @@ class TestDeleteTeamBYOKModelGhost: mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() @@ -2989,9 +2873,7 @@ class TestDeleteTeamBYOKModelGhost: mock_router = MagicMock() mock_router.model_name_to_deployment_indices = {internal_name: [0]} - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -3044,9 +2926,7 @@ class TestDeleteModelTeamAuth: mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) # The team is gone -> every team lookup returns None. @@ -3068,9 +2948,7 @@ class TestDeleteModelTeamAuth: model_id = "orphaned-byok-1" mock_prisma = self._orphaned_model_mocks(team_id, model_id) - admin_user = UserAPIKeyAuth( - user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN - ) + admin_user = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -3106,9 +2984,7 @@ class TestDeleteModelTeamAuth: model_id = "orphaned-byok-2" mock_prisma = self._orphaned_model_mocks(team_id, model_id) - non_admin = UserAPIKeyAuth( - user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER - ) + non_admin = UserAPIKeyAuth(user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -3163,9 +3039,7 @@ class TestDeleteModelTeamAuth: mock_prisma.db = MagicMock() mock_prisma.db.litellm_proxymodeltable = AsyncMock() mock_prisma.db.query_raw = AsyncMock(return_value=[]) - mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock( - return_value=db_row - ) + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) mock_prisma.db.litellm_teamtable = AsyncMock() @@ -3175,9 +3049,7 @@ class TestDeleteModelTeamAuth: # A team member who is not the team admin: rejected before the delete runs, # so the only team lookup is the single one inside the auth check. - non_admin = UserAPIKeyAuth( - user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER - ) + non_admin = UserAPIKeyAuth(user_id="someone", user_role=LitellmUserRoles.INTERNAL_USER) _PS = "litellm.proxy.proxy_server" _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" @@ -3371,15 +3243,11 @@ class TestDeleteTeamModels: prisma = _TxPrismaClient(rows) router = _RecordingRouter(prisma.events) - await delete_team_models( - team_ids=["team_a", "team_b"], prisma_client=prisma, llm_router=router - ) + await delete_team_models(team_ids=["team_a", "team_b"], prisma_client=prisma, llm_router=router) commit_idx = prisma.events.index(("commit",)) router_indices = [i for i, e in enumerate(prisma.events) if e[0] == "router"] - delete_indices = [ - i for i, e in enumerate(prisma.events) if e[0] == "delete_many" - ] + delete_indices = [i for i, e in enumerate(prisma.events) if e[0] == "delete_many"] assert router_indices, "router was never synced" assert all(i > commit_idx for i in router_indices) assert all(i < commit_idx for i in delete_indices) @@ -3395,9 +3263,7 @@ class TestDeleteTeamModels: prisma = _TxPrismaClient([mine, intruder]) router = _RecordingRouter(prisma.events) - deleted = await delete_team_models( - team_ids=["team_a"], prisma_client=prisma, llm_router=router - ) + deleted = await delete_team_models(team_ids=["team_a"], prisma_client=prisma, llm_router=router) assert deleted == ["a1"] assert router.deleted == ["a1"] @@ -3407,9 +3273,7 @@ class TestDeleteTeamModels: prisma = _TxPrismaClient([]) router = _RecordingRouter(prisma.events) - deleted = await delete_team_models( - team_ids=["team_a"], prisma_client=prisma, llm_router=router - ) + deleted = await delete_team_models(team_ids=["team_a"], prisma_client=prisma, llm_router=router) assert deleted == [] assert router.deleted == [] @@ -3420,9 +3284,7 @@ class TestDeleteTeamModels: rows = [_model_row("a1", "team_a")] prisma = _TxPrismaClient(rows) - deleted = await delete_team_models( - team_ids=["team_a"], prisma_client=prisma, llm_router=None - ) + deleted = await delete_team_models(team_ids=["team_a"], prisma_client=prisma, llm_router=None) assert deleted == ["a1"] assert any(e[0] == "delete_many" for e in prisma.events) @@ -3628,9 +3490,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(input_cost_per_token=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(input_cost_per_token=None)), ) params = json.loads(result["litellm_params"]) @@ -3649,9 +3509,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(output_cost_per_token=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(output_cost_per_token=None)), ) params = json.loads(result["litellm_params"]) @@ -3667,9 +3525,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(input_cost_per_token=0.000005) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(input_cost_per_token=0.000005)), ) params = json.loads(result["litellm_params"]) @@ -3684,9 +3540,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(output_cost_per_token=0.000007) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(output_cost_per_token=0.000007)), ) params = json.loads(result["litellm_params"]) @@ -3721,9 +3575,7 @@ class TestUpdateDBModelClearPricing: # or any other non-pricing field from the merged dict. result = update_db_model( db_model=db_model, - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(api_base=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(api_base=None)), ) info = json.loads(result["model_info"]) @@ -3758,9 +3610,7 @@ class TestUpdateDBModelClearPricing: params = json.loads(result["litellm_params"]) info = json.loads(result["model_info"]) assert "input_cost_per_token" not in params - assert ( - "input_cost_per_token" not in info - ), "model_info passthrough must not resurrect the cleared override" + assert "input_cost_per_token" not in info, "model_info passthrough must not resurrect the cleared override" def test_clear_via_model_info_clears_both_blobs(self): """The mirror works in the reverse direction too: nulling a pricing field @@ -3772,9 +3622,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=_build_db_model_with_pricing(), - updated_patch=updateDeployment( - model_info=ModelInfo(id="dep-pricing-0", input_cost_per_token=None) - ), + updated_patch=updateDeployment(model_info=ModelInfo(id="dep-pricing-0", input_cost_per_token=None)), ) params = json.loads(result["litellm_params"]) @@ -3806,9 +3654,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=db_model, - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(cache_read_input_token_cost=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(cache_read_input_token_cost=None)), ) params = json.loads(result["litellm_params"]) @@ -3840,9 +3686,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=db_model, - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(cache_creation_input_token_cost=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(cache_creation_input_token_cost=None)), ) params = json.loads(result["litellm_params"]) @@ -3876,9 +3720,7 @@ class TestUpdateDBModelClearPricing: result = update_db_model( db_model=db_model, - updated_patch=updateDeployment( - litellm_params=updateLiteLLMParams(cache_read_input_token_cost=None) - ), + updated_patch=updateDeployment(litellm_params=updateLiteLLMParams(cache_read_input_token_cost=None)), ) params = json.loads(result["litellm_params"]) @@ -8315,3 +8157,1023 @@ class TestModelManagementActorEdges: assert response.status_code == 400 assert "Cannot edit config-based model" in response.text prisma.db.litellm_proxymodeltable.update.assert_not_awaited() + + +class TestAddModelToDbBlocked: + """`_add_model_to_db` must thread `blocked` into the initial insert, so the wizard can + create a discovered-but-unchecked model already paused instead of active-then-patched.""" + + @staticmethod + def _deployment(blocked): + from litellm.types.router import ModelInfo + + return Deployment( + model_name="anthropic/claude-discovered", + litellm_params=LiteLLM_Params(model="anthropic/claude-discovered"), + model_info=ModelInfo(id="dep-blocked-create-0"), + blocked=blocked, + ) + + @pytest.mark.asyncio + async def test_add_model_to_db_writes_blocked_true(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _add_model_to_db, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=MagicMock()) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ): # test-quality-ok: the proxy wiring under test is what this patches + await _add_model_to_db( + model_params=self._deployment(True), user_api_key_dict=admin, prisma_client=mock_prisma + ) + + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert kwargs["data"]["blocked"] is True + + @pytest.mark.asyncio + async def test_add_model_to_db_writes_blocked_false(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _add_model_to_db, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=MagicMock()) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ): # test-quality-ok: the proxy wiring under test is what this patches + await _add_model_to_db( + model_params=self._deployment(False), user_api_key_dict=admin, prisma_client=mock_prisma + ) + + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert kwargs["data"]["blocked"] is False + + @pytest.mark.asyncio + async def test_add_model_to_db_omits_blocked_when_not_set(self): + """None means "don't set it" -- the Prisma column defaults to False -- not "explicitly + unblocked", so the key must be absent from the write entirely.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _add_model_to_db, + ) + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=MagicMock()) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ): # test-quality-ok: the proxy wiring under test is what this patches + await _add_model_to_db( + model_params=self._deployment(None), user_api_key_dict=admin, prisma_client=mock_prisma + ) + + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert "blocked" not in kwargs["data"] + + +class TestAddNewModelBlockedAuthGate: + """Same proxy-admin-only rule patch_model applies to `blocked` must hold at create time + too: a team admin authorized for a team-scoped model must not be able to create it already + paused out from under the proxy admin. Only a blocking value is refused, since every create + already lands unblocked and clients send the whole model shape on every create.""" + + @pytest.mark.asyncio + async def test_non_admin_cannot_set_blocked_on_create(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + mock_prisma = MagicMock() + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini"), + model_info={"id": "blocked-gate-create-0"}, + blocked=True, + ), + user_api_key_dict=non_admin, + ) + assert "proxy admin" in str(exc_info.value.message).lower() + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_non_admin_can_create_a_model_with_blocked_false(self): + """The dashboard and the SDK both send the whole model shape on create, so `blocked: false` + rides along on an ordinary team-admin create. It asks for the state the create already + lands in, and refusing the flag's presence turned every one of those creates into a 403.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + mock_prisma = MagicMock() + created_row = MagicMock() + created_row.model_id = "blocked-gate-create-2" + created_row.model_dump_json.return_value = "{}" + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=created_row) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": ["blocked-gate-create-2"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.proxy_config", + MagicMock(add_deployment=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None))), + ), + ): + result = await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini"), + model_info={"id": "blocked-gate-create-2"}, + blocked=False, + ), + user_api_key_dict=non_admin, + ) + assert result is created_row + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert kwargs["data"].get("blocked") is not True + + @pytest.mark.asyncio + async def test_proxy_admin_can_create_a_blocked_model(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + mock_prisma = MagicMock() + created_row = MagicMock() + created_row.model_id = "blocked-gate-create-1" + created_row.model_dump_json.return_value = "{}" + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=created_row) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", mock_prisma + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", True + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": ["blocked-gate-create-1"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.management_endpoints.model_management_endpoints.ModelManagementAuthChecks.can_user_make_model_call", + new=AsyncMock(return_value=None), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.proxy_config", + MagicMock(add_deployment=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None))), + ), + ): + result = await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params(model="openai/gpt-4o-mini"), + model_info={"id": "blocked-gate-create-1"}, + blocked=True, + ), + user_api_key_dict=admin, + ) + assert result is created_row + _, kwargs = mock_prisma.db.litellm_proxymodeltable.create.call_args + assert kwargs["data"]["blocked"] is True + + +class TestNonAdminCannotPersistWifFieldsOnModel: + """A server-owned Anthropic WIF field (destination, source, or secret reference) chooses + which server-side secret is read and where it is sent. A team admin who is otherwise + authorized for a team-scoped model must not be able to set one via /model/new, + /model/update, or PATCH /model/{id}/update; a proxy admin still can.""" + + @pytest.mark.asyncio + async def test_patch_model_non_admin_cannot_set_wif_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + patch_model, + ) + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + existing_row = MagicMock() + existing_row.litellm_params = {"model": "anthropic/claude-sonnet-4"} + existing_row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": "m1"}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": ["m1"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", + True, + ), + ): + with pytest.raises( + Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity" + ) as exc_info: + await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams( + anthropic_keycloak_token_url="https://attacker.example/token", + ) + ), + user_api_key_dict=non_admin, + ) + err = exc_info.value + assert getattr(err, "param", "") == "anthropic_keycloak_token_url" + mock_prisma.db.litellm_proxymodeltable.update.assert_not_called() + + @pytest.mark.asyncio + async def test_patch_model_non_admin_cannot_set_openai_wif_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + patch_model, + ) + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + existing_row = MagicMock() + existing_row.litellm_params = {"model": "openai/gpt-4o-mini"} + existing_row.model_dump.return_value = { + "model_name": "gpt", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": "m1"}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": ["m1"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", + True, + ), + ): + with pytest.raises( + Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity" + ) as exc_info: + await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams( + openai_identity_token_file="/var/run/secrets/tokens/attacker", + ) + ), + user_api_key_dict=non_admin, + ) + assert getattr(exc_info.value, "param", "") == "openai_identity_token_file" + mock_prisma.db.litellm_proxymodeltable.update.assert_not_called() + + @pytest.mark.asyncio + async def test_patch_model_admin_can_set_wif_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + patch_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + existing_row = MagicMock() + existing_row.litellm_params = {"model": "anthropic/claude-sonnet-4"} + existing_row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": "m1"}, + } + existing_row.model_dump_json.return_value = "{}" + updated_row = MagicMock() + updated_row.model_dump_json.return_value = "{}" + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": ["m1"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value: value, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), + ), + ): + result = await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams( + anthropic_keycloak_token_url="https://keycloak.internal/token", + ) + ), + user_api_key_dict=admin, + ) + assert result is updated_row + mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() + + @pytest.mark.asyncio + async def test_add_new_model_non_admin_cannot_set_wif_field(self): + from litellm.proxy._types import ProxyException + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + mock_prisma = MagicMock() + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", + True, + ), + ): + with pytest.raises(ProxyException) as exc_info: + await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params( + model="anthropic/claude-sonnet-4", + anthropic_keycloak_client_secret_ref="os.environ/LITELLM_MASTER_KEY", + ), + model_info={"id": "wif-gate-create-0"}, + ), + user_api_key_dict=non_admin, + ) + assert "proxy admin" in str(exc_info.value.message).lower() + assert exc_info.value.param == "anthropic_keycloak_client_secret_ref" + mock_prisma.db.litellm_proxymodeltable.create.assert_not_called() + + @pytest.mark.asyncio + async def test_add_new_model_admin_can_set_wif_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + add_new_model, + ) + + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + mock_prisma = MagicMock() + created_row = MagicMock() + created_row.model_id = "wif-gate-create-1" + created_row.model_dump_json.return_value = "{}" + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=created_row) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", + "sk-test-master", + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": ["wif-gate-create-1"]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.proxy_config", + MagicMock(add_deployment=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None))), + ), + ): + result = await add_new_model( + model_params=Deployment( + model_name="my-model", + litellm_params=LiteLLM_Params( + model="anthropic/claude-sonnet-4", + anthropic_keycloak_client_secret_ref="os.environ/ANTHROPIC_WIF_CLIENT_SECRET", + ), + model_info={"id": "wif-gate-create-1"}, + ), + user_api_key_dict=admin, + ) + assert result is created_row + + @pytest.mark.asyncio + async def test_update_model_non_admin_cannot_set_wif_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_model, + ) + + model_id = "wif-gate-update-0" + existing_row = MagicMock() + existing_row.litellm_params = {"model": "anthropic/claude-sonnet-4"} + existing_row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": model_id}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": [model_id]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", + True, + ), + ): + with pytest.raises( + Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity" + ) as exc_info: + await update_model( + model_params=updateDeployment( + litellm_params=updateLiteLLMParams( + anthropic_keycloak_client_secret_ref="os.environ/LITELLM_MASTER_KEY", + ), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=non_admin, + ) + assert getattr(exc_info.value, "param", "") == "anthropic_keycloak_client_secret_ref" + mock_prisma.db.litellm_proxymodeltable.update.assert_not_called() + + @pytest.mark.asyncio + async def test_update_model_admin_can_set_wif_field(self): + from litellm.proxy.management_endpoints.model_management_endpoints import ( + update_model, + ) + + model_id = "wif-gate-update-1" + existing_row = MagicMock() + existing_row.litellm_params = {"model": "anthropic/claude-sonnet-4"} + existing_row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": existing_row.litellm_params, + "model_info": {"id": model_id}, + } + existing_row.model_dump_json.return_value = "{}" + updated_row = MagicMock() + updated_row.model_dump_json.return_value = "{}" + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=updated_row) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.prisma_client", + mock_prisma, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.llm_router", + MagicMock(**{"get_model_ids.return_value": [model_id]}), + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.store_model_in_db", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.premium_user", + True, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value: value, + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), + ), + ): + await update_model( + model_params=updateDeployment( + litellm_params=updateLiteLLMParams( + anthropic_keycloak_client_secret_ref="os.environ/ANTHROPIC_WIF_CLIENT_SECRET", + ), + model_info=ModelInfo(id=model_id), + ), + user_api_key_dict=admin, + ) + mock_prisma.db.litellm_proxymodeltable.update.assert_awaited_once() + written_litellm_params = mock_prisma.db.litellm_proxymodeltable.update.await_args.kwargs["data"][ + "litellm_params" + ] + assert "anthropic_keycloak_client_secret_ref" in written_litellm_params + assert "os.environ/ANTHROPIC_WIF_CLIENT_SECRET" in written_litellm_params + + +class TestOneCredentialFeedsManyModelsNoWifCopy: + """Regression: one named WIF credential feeds multiple model rows, and no WIF field is + ever copied onto a model row -- litellm_params carries only `model` and + `litellm_credential_name`, the same shape the wizard's per-row /model/new call produces.""" + + @pytest.mark.asyncio + async def test_two_discovered_models_share_the_credential_reference_only(self): + from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper + from litellm.proxy.management_endpoints.model_management_endpoints import ( + _add_model_to_db, + ) + from litellm.types.router import ModelInfo + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.create = AsyncMock(return_value=MagicMock()) + admin = UserAPIKeyAuth(user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN) + + with ( + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.proxy_server.master_key", "sk-test-master" + ), + patch( # test-quality-ok: the proxy wiring under test is what this patches + "litellm.proxy.common_utils.encrypt_decrypt_utils._get_salt_key", return_value="sk-test-master" + ), + ): + for i, discovered_id in enumerate(["claude-a", "claude-b"]): + model_params = Deployment( + model_name=discovered_id, + litellm_params=LiteLLM_Params( + model=f"anthropic/{discovered_id}", litellm_credential_name="anthropic-wif" + ), + model_info=ModelInfo(id=f"dep-shared-{i}"), + blocked=False, + ) + await _add_model_to_db(model_params=model_params, user_api_key_dict=admin, prisma_client=mock_prisma) + + assert mock_prisma.db.litellm_proxymodeltable.create.await_count == 2 + for call in mock_prisma.db.litellm_proxymodeltable.create.await_args_list: + written_litellm_params = json.loads(call.kwargs["data"]["litellm_params"]) + decrypted_credential_name = decrypt_value_helper( + value=written_litellm_params["litellm_credential_name"], key="litellm_credential_name" + ) + assert decrypted_credential_name == "anthropic-wif" + assert "anthropic_federation_rule_id" not in written_litellm_params + assert "anthropic_identity_token" not in written_litellm_params + assert call.kwargs["data"]["blocked"] is False + + +class TestWifBoundaryReadsTheResultingDeployment: + """The proxy-admin rule has to be evaluated against the deployment the write PRODUCES. + Reading only the submitted payload let a team admin keep an existing federated deployment + and change it anyway, because the fields they sent named nothing federated.""" + + @staticmethod + def _existing_wif_row(): + row = MagicMock() + row.litellm_params = { + "model": "anthropic/claude-sonnet-4", + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + } + row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": row.litellm_params, + "model_info": {"id": "m1"}, + } + return row + + @pytest.mark.asyncio + async def test_non_admin_cannot_retarget_an_existing_wif_deployment_via_api_base(self): + """api_base is not a federation field, so the payload-only check saw nothing to refuse, + and the merged deployment then sent its assertion and minted token to the new host.""" + from litellm.proxy.management_endpoints.model_management_endpoints import patch_model + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=self._existing_wif_row()) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test + patch( # test-quality-ok: proxy wiring under test + "litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]}) + ), + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test + ): + with pytest.raises( + Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity" + ): + await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams(api_base="https://gateway.internal") + ), + user_api_key_dict=non_admin, + ) + + @pytest.mark.asyncio + async def test_non_admin_cannot_detach_a_federated_credential_to_escape_the_gate(self): + """Clearing the credential name must not be the way out. A deployment federated through a + named credential carries no federation field of its own, so a patch that sends + litellm_credential_name: null alongside an api_base of the caller's choosing would leave + nothing federated to find, and the write would be allowed.""" + from litellm.proxy.management_endpoints.model_management_endpoints import patch_model + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + federated_row = MagicMock() + federated_row.litellm_params = { + "model": "anthropic/claude-sonnet-4", + "litellm_credential_name": "admin-wif", + } + federated_row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": federated_row.litellm_params, + "model_info": {"id": "m1"}, + } + + admin_credential_row = { + "credential_name": "admin-wif", + "credential_values": { + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + }, + "credential_info": {"custom_llm_provider": "anthropic"}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=federated_row) + mock_prisma.db.litellm_credentialstable.find_unique = AsyncMock(return_value=admin_credential_row) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test + patch( # test-quality-ok: proxy wiring under test + "litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]}) + ), + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test + ): + with pytest.raises( + Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity" + ): + await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams( + litellm_credential_name=None, api_base="https://gateway.internal" + ) + ), + user_api_key_dict=non_admin, + ) + + @pytest.mark.asyncio + async def test_non_admin_cannot_attach_a_federated_credential_by_name(self): + """litellm_credential_name names no federation field itself, but request-time hydration + imports whatever the credential holds, so the resulting deployment federates.""" + from litellm.proxy.management_endpoints.model_management_endpoints import patch_model + + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + plain_row = MagicMock() + plain_row.litellm_params = {"model": "anthropic/claude-sonnet-4"} + plain_row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": plain_row.litellm_params, + "model_info": {"id": "m1"}, + } + # The credential is served from the row rather than this pod's memory, which is both the + # multi-pod case and the one the gate must not miss. + admin_credential_row = { + "credential_name": "admin-wif", + "credential_values": { + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + }, + "credential_info": {"custom_llm_provider": "anthropic"}, + } + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=plain_row) + mock_prisma.db.litellm_credentialstable.find_unique = AsyncMock(return_value=admin_credential_row) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test + patch( # test-quality-ok: proxy wiring under test + "litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]}) + ), + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test + ): + with pytest.raises( + Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity" + ): + await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams(litellm_credential_name="admin-wif") + ), + user_api_key_dict=non_admin, + ) + + @pytest.mark.asyncio + async def test_non_admin_cannot_modify_a_deployment_whose_stored_credential_name_is_encrypted(self, monkeypatch): + """Rows written through /model/new hold every litellm_params value encrypted, so a gate that + looks the stored credential name up as written asks about a ciphertext, finds no such + credential, and lets the write through.""" + from litellm.proxy.management_endpoints.model_management_endpoints import patch_model + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-1234") + non_admin = UserAPIKeyAuth(user_id="team_admin", user_role=LitellmUserRoles.INTERNAL_USER) + federated_row = MagicMock() + federated_row.litellm_params = { + "model": encrypt_value_helper(value="anthropic/claude-sonnet-4"), + "litellm_credential_name": encrypt_value_helper(value="admin-wif"), + } + assert federated_row.litellm_params["litellm_credential_name"] != "admin-wif" + federated_row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": federated_row.litellm_params, + "model_info": {"id": "m1"}, + } + admin_credential_row = { + "credential_name": "admin-wif", + "credential_values": { + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + }, + "credential_info": {"custom_llm_provider": "anthropic"}, + } + + def credential_by_exact_name(**kwargs): + return admin_credential_row if kwargs["where"].get("credential_name") == "admin-wif" else None + + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=federated_row) + mock_prisma.db.litellm_credentialstable.find_unique = AsyncMock(side_effect=credential_by_exact_name) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test + patch( # test-quality-ok: proxy wiring under test + "litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]}) + ), + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test + ): + with pytest.raises( + Exception, match="Only proxy admins can change the credentials of a deployment configured for workload identity" + ): + await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams(api_base="https://gateway.internal") + ), + user_api_key_dict=non_admin, + ) + mock_prisma.db.litellm_proxymodeltable.update.assert_not_called() + + +class TestFederationGateScopesToWhatTheWriteTouches: + """The gate reads what the write SETS, not only what the row stores. Refusing every write to a + federated deployment took rate limits, renames, tags and deletion away from the team admins who + own the model, because a proxy admin federating it once made every later team edit a 403.""" + + _TEAM_ID = "wif-scope-team" + + @staticmethod + def _team_admin(): + return UserAPIKeyAuth( + user_id="team_admin", + team_id=TestFederationGateScopesToWhatTheWriteTouches._TEAM_ID, + user_role=LitellmUserRoles.INTERNAL_USER, + ) + + @classmethod + def _federated_row(cls): + row = MagicMock() + row.litellm_params = { + "model": "anthropic/claude-sonnet-4", + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + } + row.model_dump.return_value = { + "model_name": "claude", + "litellm_params": row.litellm_params, + "model_info": {"id": "m1", "team_id": cls._TEAM_ID}, + } + row.model_dump_json.return_value = "{}" + return row + + @classmethod + def _prisma_with_live_team(cls, existing_row): + team_row = LiteLLM_TeamTable( + team_id=cls._TEAM_ID, + team_alias="wif-scope-team", + members_with_roles=[Member(user_id="team_admin", role="admin")], + ) + mock_prisma = MagicMock() + mock_prisma.db.litellm_proxymodeltable.find_unique = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_proxymodeltable.update = AsyncMock(return_value=existing_row) + mock_prisma.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_row) + return mock_prisma + + @pytest.mark.asyncio + async def test_team_admin_can_still_set_rpm_on_a_federated_deployment(self): + """rpm cannot move or re-scope the token the deployment mints, so it stays a team edit.""" + from litellm.proxy.management_endpoints.model_management_endpoints import patch_model + + existing_row = self._federated_row() + mock_prisma = self._prisma_with_live_team(existing_row) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test + patch( # test-quality-ok: proxy wiring under test + "litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]}) + ), + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test + patch( # test-quality-ok: proxy wiring under test + "litellm.proxy.management_endpoints.model_management_endpoints.encrypt_value_helper", + side_effect=lambda value: value, + ), + patch( # test-quality-ok: proxy wiring under test + "litellm.proxy.management_endpoints.model_management_endpoints.clear_cache", + new=AsyncMock(return_value=ReconcileOutcome(still_desired=None, live_after=None)), + ), + ): + result = await patch_model( + model_id="m1", + patch_data=updateDeployment(litellm_params=updateLiteLLMParams(rpm=5)), + user_api_key_dict=self._team_admin(), + ) + + assert result is existing_row + _, kwargs = mock_prisma.db.litellm_proxymodeltable.update.call_args + assert json.loads(kwargs["data"]["litellm_params"])["rpm"] == 5 + + @pytest.mark.asyncio + async def test_team_admin_still_cannot_hand_the_minted_token_to_a_clientside_override(self): + """configurable_clientside_auth_params lets a caller supply the api_base the assertion and + the token it buys are sent to, so it stays proxy-admin-only however team-owned the model is.""" + from litellm.proxy.management_endpoints.model_management_endpoints import patch_model + + mock_prisma = self._prisma_with_live_team(self._federated_row()) + + with ( + patch("litellm.proxy.proxy_server.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test + patch( # test-quality-ok: proxy wiring under test + "litellm.proxy.proxy_server.llm_router", MagicMock(**{"get_model_ids.return_value": ["m1"]}) + ), + patch("litellm.proxy.proxy_server.store_model_in_db", True), # test-quality-ok: proxy wiring under test + patch("litellm.proxy.proxy_server.premium_user", True), # test-quality-ok: proxy wiring under test + ): + with pytest.raises( + Exception, + match="Only proxy admins can change the credentials of a deployment configured for workload identity", + ): + await patch_model( + model_id="m1", + patch_data=updateDeployment( + litellm_params=updateLiteLLMParams(configurable_clientside_auth_params=["api_base"]) + ), + user_api_key_dict=self._team_admin(), + ) + mock_prisma.db.litellm_proxymodeltable.update.assert_not_called() + + @pytest.mark.asyncio + async def test_team_admin_can_still_delete_a_federated_deployment(self): + """A delete writes nothing at all, so there is no destination for it to move the token to, + and the team that owns the model must be able to take it off their page.""" + from litellm.proxy.management_endpoints.model_management_endpoints import ( + ModelInfoDelete, + delete_model as delete_model_endpoint, + ) + + db_row = LiteLLM_ProxyModelTable( + model_id="m1", + model_name="claude", + litellm_params={ + "model": "anthropic/claude-sonnet-4", + "anthropic_federation_rule_id": "fdrl_admin", + "anthropic_organization_id": "org-admin", + }, + model_info={"id": "m1", "team_id": self._TEAM_ID}, + created_by="admin", + updated_by="admin", + ) + mock_prisma = self._prisma_with_live_team(db_row) + mock_prisma.db.litellm_proxymodeltable.delete = AsyncMock(return_value=db_row) + mock_prisma.db.litellm_proxymodeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.litellm_teamtable.update = AsyncMock() + mock_prisma.db.litellm_modeltable.find_many = AsyncMock(return_value=[]) + mock_prisma.db.query_raw = AsyncMock(return_value=[]) + + _PS = "litellm.proxy.proxy_server" + _MOD = "litellm.proxy.management_endpoints.model_management_endpoints" + with ( + patch(f"{_PS}.prisma_client", mock_prisma), # test-quality-ok: proxy wiring under test + patch(f"{_PS}.store_model_in_db", True), # test-quality-ok: proxy wiring under test + patch(f"{_PS}.premium_user", True), # test-quality-ok: proxy wiring under test + patch(f"{_PS}.llm_router", MagicMock()), # test-quality-ok: proxy wiring under test + patch(f"{_PS}.proxy_logging_obj", MagicMock()), # test-quality-ok: proxy wiring under test + patch(f"{_PS}.user_api_key_cache", MagicMock()), # test-quality-ok: proxy wiring under test + patch(f"{_MOD}._refresh_cached_team", new=AsyncMock()), # test-quality-ok: proxy wiring under test + ): + result = await delete_model_endpoint( + model_info=ModelInfoDelete(id="m1"), + user_api_key_dict=self._team_admin(), + ) + + assert "deleted successfully" in result["message"] + mock_prisma.db.litellm_proxymodeltable.delete.assert_awaited_once() diff --git a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index ac010cd90a0..52ebc3a881d 100644 --- a/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/unit/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -99,36 +99,28 @@ class TestBaseOpenAIPassThroughHandler: # Test joining base URL with no path and a path base_url = httpx.URL("https://api.example.com") path = "/v1/chat/completions" - result = _join_url_paths( - base_url, path, litellm.LlmProviders.OPENAI.value - ) + result = _join_url_paths(base_url, path, litellm.LlmProviders.OPENAI.value) print(f"Base URL with no path: '{base_url}' + '{path}' → '{result}'") assert str(result) == "https://api.example.com/v1/chat/completions" # Test joining base URL with path and another path base_url = httpx.URL("https://api.example.com/v1") path = "/chat/completions" - result = _join_url_paths( - base_url, path, litellm.LlmProviders.OPENAI.value - ) + result = _join_url_paths(base_url, path, litellm.LlmProviders.OPENAI.value) print(f"Base URL with path: '{base_url}' + '{path}' → '{result}'") assert str(result) == "https://api.example.com/v1/chat/completions" # Test with path not starting with slash base_url = httpx.URL("https://api.example.com/v1") path = "chat/completions" - result = _join_url_paths( - base_url, path, litellm.LlmProviders.OPENAI.value - ) + result = _join_url_paths(base_url, path, litellm.LlmProviders.OPENAI.value) print(f"Path without leading slash: '{base_url}' + '{path}' → '{result}'") assert str(result) == "https://api.example.com/v1/chat/completions" # Test with base URL having trailing slash base_url = httpx.URL("https://api.example.com/v1/") path = "/chat/completions" - result = _join_url_paths( - base_url, path, litellm.LlmProviders.OPENAI.value - ) + result = _join_url_paths(base_url, path, litellm.LlmProviders.OPENAI.value) print(f"Base URL with trailing slash: '{base_url}' + '{path}' → '{result}'") assert str(result) == "https://api.example.com/v1/chat/completions" @@ -147,17 +139,13 @@ class TestBaseOpenAIPassThroughHandler: headers = {"authorization": "Bearer test_key"} # Test with assistants API request - result = BaseOpenAIPassThroughHandler._append_openai_beta_header( - headers, assistants_request - ) + result = BaseOpenAIPassThroughHandler._append_openai_beta_header(headers, assistants_request) print(f"Assistants API request: Added header: {result}") assert result["OpenAI-Beta"] == "assistants=v2" # Test with non-assistants API request headers = {"authorization": "Bearer test_key"} - result = BaseOpenAIPassThroughHandler._append_openai_beta_header( - headers, non_assistants_request - ) + result = BaseOpenAIPassThroughHandler._append_openai_beta_header(headers, non_assistants_request) print(f"Non-assistants API request: Headers: {result}") assert "OpenAI-Beta" not in result @@ -167,9 +155,7 @@ class TestBaseOpenAIPassThroughHandler: assistant_request.url.path = "/v1/assistants/asst_123456" headers = {"authorization": "Bearer test_key"} - result = BaseOpenAIPassThroughHandler._append_openai_beta_header( - headers, assistant_request - ) + result = BaseOpenAIPassThroughHandler._append_openai_beta_header(headers, assistant_request) print(f"Assistant API request: Added header: {result}") assert result["OpenAI-Beta"] == "assistants=v2" @@ -190,9 +176,7 @@ class TestBaseOpenAIPassThroughHandler: "test-header": "value", }, ): - result = BaseOpenAIPassThroughHandler._assemble_headers( - api_key, mock_request - ) + result = BaseOpenAIPassThroughHandler._assemble_headers(api_key, mock_request) print(f"Assembled headers: {result}") assert result["authorization"] == "Bearer test_api_key" assert result["api-key"] == "test_api_key" @@ -232,9 +216,7 @@ class TestBaseOpenAIPassThroughHandler: # Verify create_pass_through_route was called with correct parameters call_args = mock_create_pass_through.call_args[1] - print( - f"create_pass_through_route called with endpoint: {call_args['endpoint']}" - ) + print(f"create_pass_through_route called with endpoint: {call_args['endpoint']}") print(f"create_pass_through_route called with target: {call_args['target']}") assert call_args["endpoint"] == "/chat/completions" assert call_args["target"] == "https://api.openai.com/v1/chat/completions" @@ -286,9 +268,7 @@ class TestVertexAIPassThroughHandler: # Mock request mock_request = Mock() - mock_request.state = ( - None # Prevent Mock from returning a truthy _cached_headers - ) + mock_request.state = None # Prevent Mock from returning a truthy _cached_headers mock_request.method = "POST" mock_request.headers = { "Authorization": "Bearer test-creds", @@ -306,9 +286,7 @@ class TestVertexAIPassThroughHandler: test_token = vertex_credentials with ( - mock.patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth" - ) as mock_load_auth, + mock.patch("litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth") as mock_load_auth, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -391,9 +369,7 @@ class TestVertexAIPassThroughHandler: # Mock request mock_request = Mock() - mock_request.state = ( - None # Prevent Mock from returning a truthy _cached_headers - ) + mock_request.state = None # Prevent Mock from returning a truthy _cached_headers mock_request.method = "POST" mock_request.headers = { "Authorization": "Bearer test-creds", @@ -411,9 +387,7 @@ class TestVertexAIPassThroughHandler: test_token = vertex_credentials with ( - mock.patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth" - ) as mock_load_auth, + mock.patch("litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth") as mock_load_auth, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -474,9 +448,7 @@ class TestVertexAIPassThroughHandler: ], ) @pytest.mark.asyncio - async def test_vertex_passthrough_with_default_credentials( - self, monkeypatch, initial_endpoint - ): + async def test_vertex_passthrough_with_default_credentials(self, monkeypatch, initial_endpoint): """ Test that when no passthrough credentials are set, default credentials are used in the request """ @@ -515,9 +487,7 @@ class TestVertexAIPassThroughHandler: mock_response = Response() with ( - mock.patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth" - ) as mock_load_auth, + mock.patch("litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth") as mock_load_auth, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -660,17 +630,13 @@ class TestVertexAIPassThroughHandler: mock_request.method = "POST" mock_response = Mock() - with patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth" - ) as mock_auth: + with patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth") as mock_auth: mock_auth.return_value = {"api_key": "test-key-123"} with patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_pass_through: - mock_pass_through.return_value = AsyncMock( - return_value={"status": "success"} - ) + mock_pass_through.return_value = AsyncMock(return_value={"status": "success"}) with pytest.raises(HTTPException) as exc_info: await vertex_proxy_route( @@ -730,7 +696,9 @@ class TestVertexAIPassThroughHandler: mock_logging_obj.model_call_details = {} # Test URL with multimodal embedding model - url_route = "/v1/projects/test-project/locations/us-central1/publishers/google/models/multimodalembedding@001:predict" + url_route = ( + "/v1/projects/test-project/locations/us-central1/publishers/google/models/multimodalembedding@001:predict" + ) start_time = datetime.datetime.now() end_time = datetime.datetime.now() @@ -748,19 +716,13 @@ class TestVertexAIPassThroughHandler: mock_embedding_response = EmbeddingResponse( object="list", data=[ - Embedding( - embedding=[0.1, 0.2, 0.3, 0.4, 0.5], index=0, object="embedding" - ), - Embedding( - embedding=[0.6, 0.7, 0.8, 0.9, 1.0], index=1, object="embedding" - ), + Embedding(embedding=[0.1, 0.2, 0.3, 0.4, 0.5], index=0, object="embedding"), + Embedding(embedding=[0.6, 0.7, 0.8, 0.9, 1.0], index=1, object="embedding"), ], model="multimodalembedding@001", usage=Usage(prompt_tokens=0, total_tokens=0, completion_tokens=0), ) - mock_config_instance.transform_embedding_response.return_value = ( - mock_embedding_response - ) + mock_config_instance.transform_embedding_response.return_value = mock_embedding_response # Call the handler result = VertexPassthroughLoggingHandler.vertex_passthrough_handler( @@ -796,26 +758,12 @@ class TestVertexAIPassThroughHandler: ) # Test case 1: Response with textEmbedding should be detected as multimodal - response_with_text_embedding = { - "predictions": [{"textEmbedding": [0.1, 0.2, 0.3]}] - } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - response_with_text_embedding - ) - is True - ) + response_with_text_embedding = {"predictions": [{"textEmbedding": [0.1, 0.2, 0.3]}]} + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(response_with_text_embedding) is True # Test case 2: Response with imageEmbedding should be detected as multimodal - response_with_image_embedding = { - "predictions": [{"imageEmbedding": [0.4, 0.5, 0.6]}] - } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - response_with_image_embedding - ) - is True - ) + response_with_image_embedding = {"predictions": [{"imageEmbedding": [0.4, 0.5, 0.6]}]} + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(response_with_image_embedding) is True # Test case 3: Response with videoEmbeddings should be detected as multimodal response_with_video_embeddings = { @@ -831,43 +779,19 @@ class TestVertexAIPassThroughHandler: } ] } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - response_with_video_embeddings - ) - is True - ) + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(response_with_video_embeddings) is True # Test case 4: Regular text embedding response should NOT be detected as multimodal - regular_embedding_response = { - "predictions": [{"embeddings": {"values": [0.1, 0.2, 0.3]}}] - } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - regular_embedding_response - ) - is False - ) + regular_embedding_response = {"predictions": [{"embeddings": {"values": [0.1, 0.2, 0.3]}}]} + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(regular_embedding_response) is False # Test case 5: Non-embedding response should NOT be detected as multimodal - non_embedding_response = { - "candidates": [{"content": {"parts": [{"text": "Hello world"}]}}] - } - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - non_embedding_response - ) - is False - ) + non_embedding_response = {"candidates": [{"content": {"parts": [{"text": "Hello world"}]}}]} + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(non_embedding_response) is False # Test case 6: Empty response should NOT be detected as multimodal empty_response = {} - assert ( - VertexPassthroughLoggingHandler._is_multimodal_embedding_response( - empty_response - ) - is False - ) + assert VertexPassthroughLoggingHandler._is_multimodal_embedding_response(empty_response) is False def test_vertex_passthrough_handler_predict_cost_tracking(self): """ @@ -907,7 +831,9 @@ class TestVertexAIPassThroughHandler: mock_logging_obj.model_call_details = {} # Test URL with /predict endpoint - url_route = "/v1/projects/test-project/locations/us-central1/publishers/google/models/textembedding-gecko@001:predict" + url_route = ( + "/v1/projects/test-project/locations/us-central1/publishers/google/models/textembedding-gecko@001:predict" + ) start_time = datetime.datetime.now() end_time = datetime.datetime.now() @@ -977,7 +903,9 @@ class TestVertexAIPassThroughHandler: mock_logging_obj.litellm_call_id = "test-call-id-embed" mock_logging_obj.model_call_details = {} - url_route = "/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-001:embedContent" + url_route = ( + "/v1/projects/test-project/locations/us-central1/publishers/google/models/gemini-embedding-001:embedContent" + ) start_time = datetime.datetime.now() end_time = datetime.datetime.now() @@ -996,9 +924,7 @@ class TestVertexAIPassThroughHandler: ) assert result is not None - assert ( - result["result"] is not None - ), "result must not be None — logging callbacks need a non-null response" + assert result["result"] is not None, "result must not be None — logging callbacks need a non-null response" assert "kwargs" in result assert result["kwargs"].get("response_cost") == 0.0002 assert result["kwargs"].get("model") == "gemini-embedding-001" @@ -1055,9 +981,7 @@ class TestVertexAIPassThroughHandler: ) assert result is not None - assert ( - result["result"] is not None - ), "result must not be None for batchEmbedContents" + assert result["result"] is not None, "result must not be None for batchEmbedContents" assert result["kwargs"].get("response_cost") == 0.0003 assert result["kwargs"].get("model") == "gemini-embedding-001" assert result["kwargs"].get("custom_llm_provider") == "vertex_ai" @@ -1114,9 +1038,9 @@ class TestVertexAIPassThroughHandler: assert result is not None assert result["result"] is not None - assert ( - result["kwargs"].get("custom_llm_provider") == "gemini" - ), "Google AI Studio embedContent URLs must set custom_llm_provider=gemini, not vertex_ai" + assert result["kwargs"].get("custom_llm_provider") == "gemini", ( + "Google AI Studio embedContent URLs must set custom_llm_provider=gemini, not vertex_ai" + ) assert result["kwargs"].get("model") == "gemini-embedding-2-preview" mock_completion_cost.assert_called_once() @@ -1263,13 +1187,13 @@ class TestVertexAIDiscoveryPassThroughHandler: pass_through_router, ) - endpoint = f"v1/projects/{vertex_project}/locations/{vertex_location}/dataStores/default/servingConfigs/default:search" + endpoint = ( + f"v1/projects/{vertex_project}/locations/{vertex_location}/dataStores/default/servingConfigs/default:search" + ) # Mock request mock_request = Mock() - mock_request.state = ( - None # Prevent Mock from returning a truthy _cached_headers - ) + mock_request.state = None # Prevent Mock from returning a truthy _cached_headers mock_request.method = "POST" mock_request.headers = { "Authorization": "Bearer test-key", @@ -1287,9 +1211,7 @@ class TestVertexAIDiscoveryPassThroughHandler: test_token = "test-auth-token" with ( - mock.patch( - "litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth" - ) as mock_load_auth, + mock.patch("litellm.llms.vertex_ai.vertex_llm_base.VertexBase.load_auth") as mock_load_auth, mock.patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -1337,10 +1259,7 @@ class TestVertexAIDiscoveryPassThroughHandler: assert test_project in call_args[1]["target"] assert test_location in call_args[1]["target"] assert "Authorization" in call_args[1]["custom_headers"] - assert ( - call_args[1]["custom_headers"]["Authorization"] - == f"Bearer {test_token}" - ) + assert call_args[1]["custom_headers"]["Authorization"] == f"Bearer {test_token}" @pytest.mark.asyncio async def test_vertex_discovery_proxy_route_api_key_auth(self): @@ -1355,17 +1274,13 @@ class TestVertexAIDiscoveryPassThroughHandler: mock_request.method = "POST" mock_response = Mock() - with patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth" - ) as mock_auth: + with patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.user_api_key_auth") as mock_auth: mock_auth.return_value = {"api_key": "test-key-123"} with patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_pass_through: - mock_pass_through.return_value = AsyncMock( - return_value={"status": "success"} - ) + mock_pass_through.return_value = AsyncMock(return_value={"status": "success"}) with pytest.raises(HTTPException) as exc_info: await vertex_discovery_proxy_route( @@ -1461,9 +1376,7 @@ async def test_mistral_passthrough_accepts_multipart_without_json_parsing(): assert response == {"ok": True} assert captured_kwargs["is_streaming_request"] is False - assert captured_kwargs["custom_headers"] == { - "Authorization": "Bearer mistral-test-key" - } + assert captured_kwargs["custom_headers"] == {"Authorization": "Bearer mistral-test-key"} class TestBedrockLLMProxyRoute: @@ -1475,9 +1388,7 @@ class TestBedrockLLMProxyRoute: mock_user_api_key_dict = Mock() mock_request_body = {"messages": [{"role": "user", "content": "test"}]} mock_processor = Mock() - mock_processor.base_passthrough_process_llm_request = AsyncMock( - return_value="success" - ) + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") with ( patch( @@ -1489,9 +1400,10 @@ class TestBedrockLLMProxyRoute: return_value=mock_processor, ), ): - # Test application-inference-profile endpoint - endpoint = "model/arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/r742sbn2zckd/converse" + endpoint = ( + "model/arn:aws:bedrock:us-east-1:123456789012:application-inference-profile/r742sbn2zckd/converse" + ) result = await bedrock_llm_proxy_route( endpoint=endpoint, @@ -1501,9 +1413,7 @@ class TestBedrockLLMProxyRoute: ) mock_processor.base_passthrough_process_llm_request.assert_called_once() - call_kwargs = ( - mock_processor.base_passthrough_process_llm_request.call_args.kwargs - ) + call_kwargs = mock_processor.base_passthrough_process_llm_request.call_args.kwargs # For application-inference-profile, model should be "arn:aws:bedrock:us-east-1:026090525607:application-inference-profile/r742sbn2zckd" assert ( @@ -1520,9 +1430,7 @@ class TestBedrockLLMProxyRoute: mock_user_api_key_dict = Mock() mock_request_body = {"messages": [{"role": "user", "content": "test"}]} mock_processor = Mock() - mock_processor.base_passthrough_process_llm_request = AsyncMock( - return_value="success" - ) + mock_processor.base_passthrough_process_llm_request = AsyncMock(return_value="success") with ( patch( @@ -1534,7 +1442,6 @@ class TestBedrockLLMProxyRoute: return_value=mock_processor, ), ): - # Test regular model endpoint endpoint = "model/anthropic.claude-3-sonnet-20240229-v1:0/converse" @@ -1545,9 +1452,7 @@ class TestBedrockLLMProxyRoute: user_api_key_dict=mock_user_api_key_dict, ) mock_processor.base_passthrough_process_llm_request.assert_called_once() - call_kwargs = ( - mock_processor.base_passthrough_process_llm_request.call_args.kwargs - ) + call_kwargs = mock_processor.base_passthrough_process_llm_request.call_args.kwargs # For regular models, model should be just the model ID assert call_kwargs["model"] == "anthropic.claude-3-sonnet-20240229-v1:0" @@ -1570,9 +1475,7 @@ class TestBedrockLLMProxyRoute: # Create a mock httpx.Response for the error mock_error_response = Mock(spec=httpx.Response) mock_error_response.status_code = 400 - mock_error_response.aread = AsyncMock( - return_value=bedrock_error_message.encode("utf-8") - ) + mock_error_response.aread = AsyncMock(return_value=bedrock_error_message.encode("utf-8")) # Create the HTTPStatusError mock_http_error = httpx.HTTPStatusError( @@ -1589,9 +1492,7 @@ class TestBedrockLLMProxyRoute: mock_request.url = MagicMock() mock_request.url.path = "/bedrock/model/test-model/converse" - mock_request_body = { - "messages": [{"role": "user", "content": [{"textaaa": "Hello"}]}] - } + mock_request_body = {"messages": [{"role": "user", "content": [{"textaaa": "Hello"}]}]} mock_llm_router = Mock() @@ -1632,9 +1533,8 @@ class TestBedrockLLMProxyRoute: ) assert exc_info.value.status_code == 400 - assert ( - "ContentBlock object at messages.0.content.0 must set one of the following keys" - in str(exc_info.value.detail) + assert "ContentBlock object at messages.0.content.0 must set one of the following keys" in str( + exc_info.value.detail ) @pytest.mark.asyncio @@ -1712,24 +1612,14 @@ class TestBedrockLLMProxyRoute: deployment_litellm_params = deployment.get("litellm_params", {}) # Verify model-specific credentials are in the deployment - assert ( - deployment_litellm_params.get("aws_access_key_id") == model_access_key - ) - assert ( - deployment_litellm_params.get("aws_secret_access_key") - == model_secret_key - ) + assert deployment_litellm_params.get("aws_access_key_id") == model_access_key + assert deployment_litellm_params.get("aws_secret_access_key") == model_secret_key assert deployment_litellm_params.get("aws_region_name") == model_region - assert ( - deployment_litellm_params.get("aws_session_token") - == model_session_token - ) + assert deployment_litellm_params.get("aws_session_token") == model_session_token # Verify environment variables are NOT in the deployment assert deployment_litellm_params.get("aws_access_key_id") != env_access_key - assert ( - deployment_litellm_params.get("aws_secret_access_key") != env_secret_key - ) + assert deployment_litellm_params.get("aws_secret_access_key") != env_secret_key assert deployment_litellm_params.get("aws_region_name") != env_region # Test 3: Verify credentials are passed through the passthrough route @@ -1740,9 +1630,7 @@ class TestBedrockLLMProxyRoute: captured_kwargs.update(kwargs) mock_response = MagicMock() mock_response.status_code = 200 - mock_response.aread = AsyncMock( - return_value=b'{"content": [{"text": "Hello"}]}' - ) + mock_response.aread = AsyncMock(return_value=b'{"content": [{"text": "Hello"}]}') return mock_response mock_request = MagicMock(spec=Request) @@ -1752,9 +1640,7 @@ class TestBedrockLLMProxyRoute: mock_request.url = MagicMock() mock_request.url.path = "/bedrock/model/claude-opus-4-1/converse" - mock_request_body = { - "messages": [{"role": "user", "content": [{"text": "Hello"}]}] - } + mock_request_body = {"messages": [{"role": "user", "content": [{"text": "Hello"}]}]} mock_user_api_key_dict = Mock() mock_user_api_key_dict.api_key = "test-key" @@ -1775,9 +1661,7 @@ class TestBedrockLLMProxyRoute: # Setup mock response mock_response = MagicMock() mock_response.status_code = 200 - mock_response.aread = AsyncMock( - return_value=b'{"content": [{"text": "Hello"}]}' - ) + mock_response.aread = AsyncMock(return_value=b'{"content": [{"text": "Hello"}]}') mock_process.return_value = mock_response # Call the handler @@ -2143,9 +2027,7 @@ class TestLLMPassthroughFactoryProxyRoute: mock_user_api_key_dict = MagicMock() with ( - patch( - "litellm.utils.ProviderConfigManager.get_provider_model_info" - ) as mock_get_provider, + patch("litellm.utils.ProviderConfigManager.get_provider_model_info") as mock_get_provider, patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials" ) as mock_get_creds, @@ -2155,9 +2037,7 @@ class TestLLMPassthroughFactoryProxyRoute: ): mock_provider_config = MagicMock() mock_provider_config.get_api_base.return_value = "https://example.com/v1" - mock_provider_config.validate_environment.return_value = { - "x-api-key": "dummy" - } + mock_provider_config.validate_environment.return_value = {"x-api-key": "dummy"} mock_get_provider.return_value = mock_provider_config mock_get_creds.return_value = "dummy" @@ -2173,12 +2053,8 @@ class TestLLMPassthroughFactoryProxyRoute: ) assert result == "success" - mock_get_provider.assert_called_once_with( - provider=litellm.LlmProviders(LlmProviders.VLLM), model=None - ) - mock_get_creds.assert_called_once_with( - custom_llm_provider=LlmProviders.VLLM, region_name=None - ) + mock_get_provider.assert_called_once_with(provider=litellm.LlmProviders(LlmProviders.VLLM), model=None) + mock_get_creds.assert_called_once_with(custom_llm_provider=LlmProviders.VLLM, region_name=None) mock_create_route.assert_called_once_with( endpoint="/chat/completions", target="https://example.com/v1/chat/completions", @@ -2576,9 +2452,7 @@ class TestForwardHeaders: # Create a mock request with custom headers mock_request = MagicMock(spec=Request) - mock_request.state = ( - None # Prevent MagicMock from returning a truthy _cached_headers - ) + mock_request.state = None # Prevent MagicMock from returning a truthy _cached_headers mock_request.method = "POST" mock_request.url = MagicMock() mock_request.url.path = "/test/endpoint" @@ -2613,9 +2487,7 @@ class TestForwardHeaders: mock_httpx_response = MagicMock() mock_httpx_response.status_code = 200 mock_httpx_response.headers = {"content-type": "application/json"} - mock_httpx_response.aiter_bytes = AsyncMock( - return_value=[b'{"result": "success"}'] - ) + mock_httpx_response.aiter_bytes = AsyncMock(return_value=[b'{"result": "success"}']) mock_httpx_response.aread = AsyncMock(return_value=b'{"result": "success"}') with ( @@ -2640,9 +2512,7 @@ class TestForwardHeaders: mock_logging_obj.pre_call_hook = AsyncMock(return_value=mock_request_body) mock_logging_obj.post_call_success_hook = AsyncMock() mock_logging_obj.post_call_failure_hook = AsyncMock() - mock_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value={} - ) + mock_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={}) # Call pass_through_request with forward_headers=True result = await pass_through_request( @@ -2715,9 +2585,7 @@ class TestForwardHeaders: mock_httpx_response = MagicMock() mock_httpx_response.status_code = 200 mock_httpx_response.headers = {"content-type": "application/json"} - mock_httpx_response.aiter_bytes = AsyncMock( - return_value=[b'{"result": "success"}'] - ) + mock_httpx_response.aiter_bytes = AsyncMock(return_value=[b'{"result": "success"}']) mock_httpx_response.aread = AsyncMock(return_value=b'{"result": "success"}') with ( @@ -2742,9 +2610,7 @@ class TestForwardHeaders: mock_logging_obj.pre_call_hook = AsyncMock(return_value=mock_request_body) mock_logging_obj.post_call_success_hook = AsyncMock() mock_logging_obj.post_call_failure_hook = AsyncMock() - mock_logging_obj.post_call_response_headers_hook = AsyncMock( - return_value={} - ) + mock_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={}) # Call pass_through_request with forward_headers=False (default) result = await pass_through_request( @@ -2802,15 +2668,11 @@ class TestForwardHeaders: mock_httpx_response = MagicMock() mock_httpx_response.status_code = 200 mock_httpx_response.headers = {"content-type": "application/json"} - mock_httpx_response.aiter_bytes = AsyncMock( - return_value=[b'{"result": "success"}'] - ) + mock_httpx_response.aiter_bytes = AsyncMock(return_value=[b'{"result": "success"}']) mock_httpx_response.aread = AsyncMock(return_value=b'{"result": "success"}') with ( - patch( - "litellm.utils.ProviderConfigManager.get_provider_model_info" - ) as mock_get_provider, + patch("litellm.utils.ProviderConfigManager.get_provider_model_info") as mock_get_provider, patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials" ) as mock_get_creds, @@ -2826,9 +2688,7 @@ class TestForwardHeaders: # Setup provider config mock_provider_config = MagicMock() mock_provider_config.get_api_base.return_value = "https://api.openai.com/v1" - mock_provider_config.validate_environment.return_value = { - "authorization": "Bearer sk-test" - } + mock_provider_config.validate_environment.return_value = {"authorization": "Bearer sk-test"} mock_get_provider.return_value = mock_provider_config mock_get_creds.return_value = "sk-test" @@ -2840,9 +2700,7 @@ class TestForwardHeaders: mock_get_client.return_value = mock_client_obj # Setup mock logging object - mock_logging_obj.pre_call_hook = AsyncMock( - return_value={"messages": [{"role": "user", "content": "test"}]} - ) + mock_logging_obj.pre_call_hook = AsyncMock(return_value={"messages": [{"role": "user", "content": "test"}]}) mock_logging_obj.post_call_success_hook = AsyncMock() # This is the key part - when create_pass_through_route is called with _forward_headers=True @@ -2933,24 +2791,16 @@ class TestMilvusProxyRoute: ): # Setup mocks mock_provider_config = MagicMock() - mock_provider_config.get_auth_credentials.return_value = { - "headers": {"Authorization": "Bearer test-token"} - } + mock_provider_config.get_auth_credentials.return_value = {"headers": {"Authorization": "Bearer test-token"}} mock_provider_config.get_complete_url.return_value = api_base mock_get_config.return_value = mock_provider_config mock_index_registry.is_vector_store_index.return_value = True - mock_index_registry.get_vector_store_index_by_name.return_value = ( - mock_index_object - ) + mock_index_registry.get_vector_store_index_by_name.return_value = mock_index_object - mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = ( - mock_vector_store - ) + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = mock_vector_store - mock_endpoint_func = AsyncMock( - return_value={"results": [{"id": 1, "distance": 0.5}]} - ) + mock_endpoint_func = AsyncMock(return_value={"results": [{"id": 1, "distance": 0.5}]}) mock_create_route.return_value = mock_endpoint_func # Call the route @@ -2963,9 +2813,7 @@ class TestMilvusProxyRoute: # Verify calls mock_get_body.assert_called_once() - mock_index_registry.is_vector_store_index.assert_called_once_with( - vector_store_index_name=collection_name - ) + mock_index_registry.is_vector_store_index.assert_called_once_with(vector_store_index_name=collection_name) mock_is_allowed.assert_called_once() mock_safe_set.assert_called_once() @@ -2977,9 +2825,7 @@ class TestMilvusProxyRoute: mock_create_route.assert_called_once() create_route_args = mock_create_route.call_args[1] assert "vectors/search" in create_route_args["target"] - assert create_route_args["custom_headers"] == { - "Authorization": "Bearer test-token" - } + assert create_route_args["custom_headers"] == {"Authorization": "Bearer test-token"} # Verify endpoint function was called mock_endpoint_func.assert_awaited_once() @@ -2992,7 +2838,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - mock_request = MagicMock(spec=Request) mock_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -3026,7 +2871,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - mock_request = MagicMock(spec=Request) mock_response = MagicMock(spec=Response) mock_user_api_key_dict = MagicMock() @@ -3044,9 +2888,7 @@ class TestMilvusProxyRoute: ) assert exc_info.value.status_code == 500 - assert "Unable to find Milvus vector store config" in str( - exc_info.value.detail - ) + assert "Unable to find Milvus vector store config" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_milvus_proxy_route_no_index_registry(self): @@ -3055,7 +2897,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - collection_name = "test-collection" mock_request = MagicMock(spec=Request) @@ -3083,9 +2924,7 @@ class TestMilvusProxyRoute: ) assert exc_info.value.status_code == 500 - assert "Unable to find Milvus vector store index registry" in str( - exc_info.value.detail - ) + assert "Unable to find Milvus vector store index registry" in str(exc_info.value.detail) @pytest.mark.asyncio async def test_milvus_proxy_route_not_managed_index(self): @@ -3094,7 +2933,6 @@ class TestMilvusProxyRoute: """ from fastapi import HTTPException - collection_name = "unmanaged-collection" mock_request = MagicMock(spec=Request) @@ -3124,9 +2962,8 @@ class TestMilvusProxyRoute: ) assert exc_info.value.status_code == 400 - assert ( - f"Collection {collection_name} is not a litellm managed vector store index" - in str(exc_info.value.detail) + assert f"Collection {collection_name} is not a litellm managed vector store index" in str( + exc_info.value.detail ) @pytest.mark.asyncio @@ -3158,22 +2995,16 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body" - ), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, patch.object(litellm, "vector_store_registry") as mock_vector_registry, ): mock_get_config.return_value = MagicMock() mock_index_registry.is_vector_store_index.return_value = True - mock_index_registry.get_vector_store_index_by_name.return_value = ( - mock_index_object - ) - mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = ( - None - ) + mock_index_registry.get_vector_store_index_by_name.return_value = mock_index_object + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = None - with pytest.raises(Exception, match='Vector store not found for missing-store') as exc_info: + with pytest.raises(Exception, match="Vector store not found for missing-store") as exc_info: await milvus_proxy_route( endpoint="vectors/search", request=mock_request, @@ -3181,9 +3012,7 @@ class TestMilvusProxyRoute: user_api_key_dict=mock_user_api_key_dict, ) - assert f"Vector store not found for {vector_store_name}" in str( - exc_info.value - ) + assert f"Vector store not found for {vector_store_name}" in str(exc_info.value) @pytest.mark.asyncio async def test_milvus_proxy_route_no_api_base(self): @@ -3216,9 +3045,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body" - ), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), patch.object(litellm, "vector_store_index_registry") as mock_index_registry, patch.object(litellm, "vector_store_registry") as mock_vector_registry, ): @@ -3228,14 +3055,10 @@ class TestMilvusProxyRoute: mock_get_config.return_value = mock_provider_config mock_index_registry.is_vector_store_index.return_value = True - mock_index_registry.get_vector_store_index_by_name.return_value = ( - mock_index_object - ) - mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = ( - mock_vector_store - ) + mock_index_registry.get_vector_store_index_by_name.return_value = mock_index_object + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = mock_vector_store - with pytest.raises(Exception, match='api_base not found in vector store configuration for') as exc_info: + with pytest.raises(Exception, match="api_base not found in vector store configuration for") as exc_info: await milvus_proxy_route( endpoint="vectors/search", request=mock_request, @@ -3243,10 +3066,7 @@ class TestMilvusProxyRoute: user_api_key_dict=mock_user_api_key_dict, ) - assert ( - f"api_base not found in vector store configuration for {vector_store_name}" - in str(exc_info.value) - ) + assert f"api_base not found in vector store configuration for {vector_store_name}" in str(exc_info.value) @pytest.mark.asyncio async def test_milvus_proxy_route_endpoint_without_leading_slash(self): @@ -3280,9 +3100,7 @@ class TestMilvusProxyRoute: patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.is_allowed_to_call_vector_store_endpoint" ), - patch( - "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body" - ), + patch("litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints._safe_set_request_parsed_body"), patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, @@ -3295,12 +3113,8 @@ class TestMilvusProxyRoute: mock_get_config.return_value = mock_provider_config mock_index_registry.is_vector_store_index.return_value = True - mock_index_registry.get_vector_store_index_by_name.return_value = ( - mock_index_object - ) - mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = ( - mock_vector_store - ) + mock_index_registry.get_vector_store_index_by_name.return_value = mock_index_object + mock_vector_registry.get_litellm_managed_vector_store_from_registry_by_name.return_value = mock_vector_store mock_endpoint_func = AsyncMock(return_value={"status": "success"}) mock_create_route.return_value = mock_endpoint_func @@ -3348,9 +3162,7 @@ class TestOpenAIPassthroughRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, ): - mock_endpoint_func = AsyncMock( - return_value={"id": "resp_123", "status": "completed"} - ) + mock_endpoint_func = AsyncMock(return_value={"id": "resp_123", "status": "completed"}) mock_create_route.return_value = mock_endpoint_func # Call the route with /v1/responses endpoint @@ -3398,9 +3210,7 @@ class TestOpenAIPassthroughRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, ): - mock_endpoint_func = AsyncMock( - return_value={"id": "chatcmpl-123", "choices": []} - ) + mock_endpoint_func = AsyncMock(return_value={"id": "chatcmpl-123", "choices": []}) mock_create_route.return_value = mock_endpoint_func result = await openai_proxy_route( @@ -3466,9 +3276,7 @@ class TestOpenAIPassthroughRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, ): - mock_endpoint_func = AsyncMock( - return_value={"id": "asst_123", "object": "assistant"} - ) + mock_endpoint_func = AsyncMock(return_value={"id": "asst_123", "object": "assistant"}) mock_create_route.return_value = mock_endpoint_func result = await openai_proxy_route( @@ -3677,9 +3485,7 @@ class TestCursorProxyRoute: "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route" ) as mock_create_route, ): - mock_endpoint_func = AsyncMock( - return_value={"agents": [], "nextCursor": None} - ) + mock_endpoint_func = AsyncMock(return_value={"agents": [], "nextCursor": None}) mock_create_route.return_value = mock_endpoint_func result = await cursor_proxy_route( @@ -3693,12 +3499,8 @@ class TestCursorProxyRoute: call_args = mock_create_route.call_args[1] assert call_args["target"] == "https://api.cursor.com/v0/agents" - expected_auth = base64.b64encode(f"{test_api_key}:".encode("utf-8")).decode( - "ascii" - ) - assert ( - call_args["custom_headers"]["Authorization"] == f"Basic {expected_auth}" - ) + expected_auth = base64.b64encode(f"{test_api_key}:".encode("utf-8")).decode("ascii") + assert call_args["custom_headers"]["Authorization"] == f"Basic {expected_auth}" assert result == {"agents": [], "nextCursor": None} @@ -3722,7 +3524,7 @@ class TestCursorProxyRoute: [], ), ): - with pytest.raises(Exception, match='Cursor API key not found\\. Add Cursor credentials via') as exc_info: + with pytest.raises(Exception, match="Cursor API key not found\\. Add Cursor credentials via") as exc_info: await cursor_proxy_route( endpoint="v0/agents", request=mock_request, @@ -3781,9 +3583,7 @@ class TestCursorProxyRoute: import base64 expected_auth = base64.b64encode(b"crsr_ui_test_key:").decode("ascii") - assert ( - call_args["custom_headers"]["Authorization"] == f"Basic {expected_auth}" - ) + assert call_args["custom_headers"]["Authorization"] == f"Basic {expected_auth}" @pytest.mark.asyncio async def test_cursor_proxy_route_custom_api_base(self): @@ -3796,9 +3596,7 @@ class TestCursorProxyRoute: mock_user_api_key_dict = MagicMock() with ( - patch.dict( - os.environ, {"CURSOR_API_BASE": "https://custom-cursor.example.com"} - ), + patch.dict(os.environ, {"CURSOR_API_BASE": "https://custom-cursor.example.com"}), patch( "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials", return_value="test-key", @@ -3878,12 +3676,10 @@ class TestVertexRawPredictStreamingClassification: """ RAW_PREDICT_ENDPOINT = ( - "v1/projects/test-project/locations/us-east5/publishers/anthropic/models/" - "claude-sonnet-4-6:streamRawPredict" + "v1/projects/test-project/locations/us-east5/publishers/anthropic/models/claude-sonnet-4-6:streamRawPredict" ) GENERATE_CONTENT_ENDPOINT = ( - "v1/projects/test-project/locations/us-east5/publishers/google/models/" - "gemini-2.5-flash:streamGenerateContent" + "v1/projects/test-project/locations/us-east5/publishers/google/models/gemini-2.5-flash:streamGenerateContent" ) async def _capture_passthrough_kwargs(self, endpoint: str, body: object) -> dict: @@ -4068,10 +3864,7 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: """ VKEY = "sk-litellm-victim-key" - ENDPOINT = ( - "v1/projects/my-proj/locations/us-central1/publishers/google/models/" - "gemini-2.5-flash:generateContent" - ) + ENDPOINT = "v1/projects/my-proj/locations/us-central1/publishers/google/models/gemini-2.5-flash:generateContent" async def _run( self, @@ -4199,7 +3992,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: (b"content-type", b"application/json"), ], ) - assert forwarded is None, f"a virtual key echoed as '{scheme} ' in Authorization must be stripped, not forwarded" + assert forwarded is None, ( + f"a virtual key echoed as '{scheme} ' in Authorization must be stripped, not forwarded" + ) assert raised is not None and raised.status_code == 401 @pytest.mark.asyncio @@ -4245,8 +4040,7 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: @pytest.mark.parametrize( "credential_header", sorted( - SpecialHeaders.litellm_credential_header_names() - - {"authorization", "x-goog-api-key", "x-litellm-api-key"} + SpecialHeaders.litellm_credential_header_names() - {"authorization", "x-goog-api-key", "x-litellm-api-key"} ), ) async def test_every_non_google_credential_header_is_dropped_by_name(self, monkeypatch, credential_header): @@ -4321,7 +4115,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: assert raised is not None and raised.status_code == 401 @pytest.mark.asyncio - async def test_authenticated_authorization_is_stripped_over_a_lower_precedence_pass_through_header(self, monkeypatch): + async def test_authenticated_authorization_is_stripped_over_a_lower_precedence_pass_through_header( + self, monkeypatch + ): with mock.patch.dict( # test-quality-ok: general_settings is the real proxy config surface for pass_through_endpoints; no injection seam exists on this route "litellm.proxy.proxy_server.general_settings", {"pass_through_endpoints": [{"headers": {"litellm_user_api_key": "x-company-key"}}]}, @@ -4338,7 +4134,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: assert raised is None assert forwarded is not None assert forwarded.get("x-goog-api-key") == "AIza-real-google-api-key" - assert "authorization" not in forwarded, "Authorization authenticated (higher precedence) so its key must be stripped" + assert "authorization" not in forwarded, ( + "Authorization authenticated (higher precedence) so its key must be stripped" + ) assert "x-company-key" not in forwarded assert self.VKEY not in " ".join(f"{name}:{value}" for name, value in forwarded.items()) @@ -4369,7 +4167,9 @@ class TestVertexCredentiallessPassthroughVirtualKeyLeak: (b"content-type", b"application/json"), ], ) - assert forwarded is None, "a virtual key in the mapped-route litellm_user_api_key header must be dropped, not forwarded" + assert forwarded is None, ( + "a virtual key in the mapped-route litellm_user_api_key header must be dropped, not forwarded" + ) assert raised is not None and raised.status_code == 401 GOOGLE_OAUTH_TOKEN = "ya29.byo-google-oauth-token" @@ -5851,6 +5651,348 @@ class TestVertexAILiveWebsocketPassthrough: assert len(close_kwargs["reason"].encode("utf-8")) <= 123 +class TestAnthropicProxyRoute: + """The /anthropic passthrough route: custom auth headers must not clobber the + client's anthropic-beta, and the WIF tier must mint through the async facade.""" + + def _get_request(self, headers: dict) -> MagicMock: + request = MagicMock(spec=Request) + request.method = "GET" + request.headers = headers + request.query_params = {} + return request + + def _clear_anthropic_env(self, monkeypatch) -> None: + for name in ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + "ANTHROPIC_FEDERATION_RULE_ID", + "ANTHROPIC_ORGANIZATION_ID", + "ANTHROPIC_IDENTITY_TOKEN_FILE", + "ANTHROPIC_IDENTITY_TOKEN", + ): + monkeypatch.delenv(name, raising=False) + + @pytest.mark.asyncio + async def test_client_anthropic_beta_merged_into_auth_header(self, monkeypatch): + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + anthropic_proxy_route, + ) + + self._clear_anthropic_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-oat01-passthrough-token") + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route", + return_value=AsyncMock(return_value={"ok": True}), + ) as mock_create_route: + await anthropic_proxy_route( + endpoint="v1/models", + request=self._get_request({"anthropic-beta": "context-1m-2025-08-07"}), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-caller-virtual-key"), + ) + + custom_headers = mock_create_route.call_args.kwargs["custom_headers"] + assert custom_headers["authorization"] == "Bearer sk-ant-oat01-passthrough-token" + betas = set(custom_headers["anthropic-beta"].split(",")) + assert {"context-1m-2025-08-07", "oauth-2025-04-20"} <= betas + + @pytest.mark.asyncio + async def test_wif_mint_goes_through_async_facade(self, monkeypatch): + import threading + + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import aget_anthropic_wif_token, get_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + anthropic_proxy_route, + ) + + self._clear_anthropic_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_route") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-route") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "route-inline-jwt") + + minted: Final = "sk-ant-oat01-route-minted" + thread_ids: Final = [] + + class ThreadRecordingPoster: + def post(self, url, *, content, headers, timeout): + thread_ids.append(threading.get_ident()) + return httpx.Response( + 200, + json={"access_token": minted, "token_type": "Bearer", "expires_in": 3600}, + ) + + engine = JwtBearerTokenExchangeEngine(poster=ThreadRecordingPoster()) + sync_calls: Final = [] + + def sync_shim(litellm_params, api_base, model): + sync_calls.append(model) + return get_anthropic_wif_token(litellm_params, api_base, model, engine) + + async def async_shim(litellm_params, api_base, model): + return await aget_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "get_anthropic_wif_token", sync_shim) + monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim) + + with patch( + "litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route", + return_value=AsyncMock(return_value={"ok": True}), + ) as mock_create_route: + await anthropic_proxy_route( + endpoint="v1/models", + request=self._get_request({}), + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-caller-virtual-key"), + ) + + custom_headers = mock_create_route.call_args.kwargs["custom_headers"] + assert custom_headers["authorization"] == f"Bearer {minted}" + assert custom_headers["anthropic-beta"] == "oauth-2025-04-20" + assert sync_calls == [] + assert thread_ids and thread_ids[0] != threading.get_ident() + + +class TestAnthropicProxyRouteCallerAuthHeaders: + """Regression for a caller credential riding upstream next to a server-owned one. + + /anthropic forwards the caller's headers, so a caller-supplied ``x-api-key`` used to reach + Anthropic alongside the server-minted ``Authorization: Bearer``. These drive the real relay + (only the httpx client is stubbed) and assert on the bytes actually handed to the upstream. + """ + + _MINTED: Final = "sk-ant-oat01-plan-minted" + + def _clear_anthropic_env(self, monkeypatch) -> None: + for name in ( + "ANTHROPIC_API_KEY", + "ANTHROPIC_AUTH_TOKEN", + "ANTHROPIC_API_BASE", + "ANTHROPIC_BASE_URL", + "ANTHROPIC_FEDERATION_RULE_ID", + "ANTHROPIC_ORGANIZATION_ID", + "ANTHROPIC_IDENTITY_TOKEN_FILE", + "ANTHROPIC_IDENTITY_TOKEN", + ): + monkeypatch.delenv(name, raising=False) + + # A sibling test leaving SERVER_ROOT_PATH set re-prefixes the passthrough route, so + # /anthropic/... stops resolving and the request 404s before any header is built. + # Pin it so this class asserts on headers rather than on ambient state. + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + + def _enable_wif(self, monkeypatch) -> None: + from litellm.llms.anthropic import common_utils as anthropic_common_utils + from litellm.llms.anthropic.wif import aget_anthropic_wif_token + from litellm.llms.base_llm.auth.token_exchange import JwtBearerTokenExchangeEngine + + monkeypatch.setenv("ANTHROPIC_FEDERATION_RULE_ID", "fdrl_plan") + monkeypatch.setenv("ANTHROPIC_ORGANIZATION_ID", "org-plan") + monkeypatch.setenv("ANTHROPIC_IDENTITY_TOKEN", "plan-inline-jwt") + + minted: Final = self._MINTED + + class StubPoster: + def post(self, url, *, content, headers, timeout): + return httpx.Response( + 200, + json={"access_token": minted, "token_type": "Bearer", "expires_in": 3600}, + ) + + engine: Final = JwtBearerTokenExchangeEngine(poster=StubPoster()) + + async def async_shim(litellm_params, api_base, model): + return await aget_anthropic_wif_token(litellm_params, api_base, model, engine) + + monkeypatch.setattr(anthropic_common_utils, "aget_anthropic_wif_token", async_shim) + + def _request(self, headers: Mapping[str, str]) -> Request: + body: Final = b'{"model":"claude-sonnet-4-5","messages":[]}' + scope: Final = { + "type": "http", + "http_version": "1.1", + "method": "POST", + "scheme": "https", + "path": "/anthropic/v1/messages", + "raw_path": b"/anthropic/v1/messages", + "root_path": "", + "query_string": b"", + "headers": [(name.lower().encode(), value.encode()) for name, value in headers.items()], + "client": ("127.0.0.1", 51234), + "server": ("proxy.local", 4000), + "state": {}, + } + + async def receive() -> dict: + return {"type": "http.request", "body": body, "more_body": False} + + return Request(scope, receive) + + async def _upstream_headers(self, request: Request) -> dict: + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + anthropic_proxy_route, + ) + + upstream_response: Final = MagicMock() + upstream_response.status_code = 200 + upstream_response.headers = {"content-type": "application/json"} + upstream_response.aread = AsyncMock(return_value=b'{"ok": true}') + upstream_response.aiter_bytes = AsyncMock(return_value=[b'{"ok": true}']) + + httpx_client: Final = MagicMock() + httpx_client.build_request = MagicMock(return_value=MagicMock()) + httpx_client.send = AsyncMock(return_value=upstream_response) + client_wrapper: Final = MagicMock() + client_wrapper.client = httpx_client + + with ( + patch( # test-quality-ok: stubbing the http client IS the boundary; the test asserts on the bytes handed to it + "litellm.proxy.pass_through_endpoints.pass_through_endpoints.get_async_httpx_client", + return_value=client_wrapper, + ), + patch( # test-quality-ok: the relay calls these hooks, and they need a db this test has no use for + "litellm.proxy.proxy_server.proxy_logging_obj" + ) as mock_logging_obj, + ): + mock_logging_obj.pre_call_hook = AsyncMock(return_value={"model": "claude-sonnet-4-5", "messages": []}) + mock_logging_obj.post_call_success_hook = AsyncMock() + mock_logging_obj.post_call_failure_hook = AsyncMock() + mock_logging_obj.post_call_response_headers_hook = AsyncMock(return_value={}) + + await anthropic_proxy_route( + endpoint="v1/messages", + request=request, + fastapi_response=MagicMock(spec=Response), + user_api_key_dict=UserAPIKeyAuth(api_key="sk-caller-virtual-key"), + ) + + assert httpx_client.send.called + return {name.lower(): value for name, value in dict(httpx_client.build_request.call_args[1]["headers"]).items()} + + @pytest.mark.asyncio + async def test_wif_credential_drops_caller_supplied_api_key(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + self._enable_wif(monkeypatch) + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "x-api-key": "sk-caller-virtual-key", + "user-agent": "caller/1.0", + } + ) + ) + + assert sent["authorization"] == f"Bearer {self._MINTED}" + assert "x-api-key" not in sent + assert sent["user-agent"] == "caller/1.0" + + @pytest.mark.asyncio + async def test_wif_credential_drops_caller_supplied_authorization(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + self._enable_wif(monkeypatch) + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "authorization": "Bearer sk-caller-virtual-key", + } + ) + ) + + assert sent["authorization"] == f"Bearer {self._MINTED}" + assert all("sk-caller-virtual-key" not in value for value in sent.values()) + + @pytest.mark.asyncio + @pytest.mark.parametrize("header_name", sorted(SpecialHeaders.litellm_credential_header_names())) + async def test_wif_credential_drops_every_proxy_key_header(self, monkeypatch, header_name: str): + """The proxy accepts a LiteLLM key in any SpecialHeaders slot, so the caller's virtual + key must not reach Anthropic from any of them once the server owns the credential.""" + self._clear_anthropic_env(monkeypatch) + self._enable_wif(monkeypatch) + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + header_name: "sk-caller-virtual-key", + "user-agent": "caller/1.0", + } + ) + ) + + assert sent["authorization"] == f"Bearer {self._MINTED}" + assert header_name == "authorization" or header_name not in sent + assert all("sk-caller-virtual-key" not in value for value in sent.values()) + assert sent["user-agent"] == "caller/1.0" + + @pytest.mark.asyncio + async def test_wif_credential_drops_configured_custom_key_header(self, monkeypatch): + from litellm.proxy import proxy_server + + self._clear_anthropic_env(monkeypatch) + self._enable_wif(monkeypatch) + monkeypatch.setitem(proxy_server.general_settings, "litellm_key_header_name", "X-Tenant-Key") + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "x-tenant-key": "sk-caller-virtual-key", + "x-tenant-region": "eu", + } + ) + ) + + assert sent["authorization"] == f"Bearer {self._MINTED}" + assert "x-tenant-key" not in sent + assert all("sk-caller-virtual-key" not in value for value in sent.values()) + assert sent["x-tenant-region"] == "eu" + + @pytest.mark.asyncio + async def test_server_api_key_drops_caller_supplied_authorization(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + monkeypatch.setenv("ANTHROPIC_API_KEY", "sk-ant-server-owned") + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "authorization": "Bearer sk-caller-virtual-key", + "x-api-key": "sk-caller-virtual-key", + } + ) + ) + + assert sent["x-api-key"] == "sk-ant-server-owned" + assert "authorization" not in sent + + @pytest.mark.asyncio + async def test_byok_caller_key_still_reaches_upstream(self, monkeypatch): + self._clear_anthropic_env(monkeypatch) + + sent: Final = await self._upstream_headers( + self._request( + { + "content-type": "application/json", + "x-api-key": "sk-ant-caller-owned", + "anthropic-version": "2023-06-01", + } + ) + ) + + assert sent["x-api-key"] == "sk-ant-caller-owned" + assert sent["anthropic-version"] == "2023-06-01" + assert "authorization" not in sent + + class TestPassthroughRouterModelBudgetReservation: """ Router-model passthrough on /vllm and /azure must thread the calling key's diff --git a/tests/unit/proxy/proxy_server/test_proxy_config.py b/tests/unit/proxy/proxy_server/test_proxy_config.py index d309e6de3b6..61aed5c9c81 100644 --- a/tests/unit/proxy/proxy_server/test_proxy_config.py +++ b/tests/unit/proxy/proxy_server/test_proxy_config.py @@ -3326,6 +3326,23 @@ async def test_ProxyConfig_load_config_warns_and_turns_off_a_non_flag_litellm_se # --------------------------------------------------------------------------- +def test_ProxyConfig_decrypt_credentials_returns_an_encrypted_empty_value_as_empty(monkeypatch): + from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper + + monkeypatch.setenv("LITELLM_SALT_KEY", "sk-decrypt-credentials-test-salt") + decrypted = ProxyConfig().decrypt_credentials( + { + "credential_name": "openai-wif", + "credential_values": { + "api_base": encrypt_value_helper(""), + "openai_service_account_id": encrypt_value_helper("user-1"), + }, + "credential_info": {"custom_llm_provider": "openai"}, + } + ) + assert decrypted.credential_values == {"api_base": "", "openai_service_account_id": "user-1"} + + def test_ProxyConfig_decrypt_model_list_from_db_returns_decrypted(monkeypatch): monkeypatch.setattr( "litellm.proxy.proxy_server.decrypt_value_helper", diff --git a/tests/unit/proxy/test_credential_slot_registry.py b/tests/unit/proxy/test_credential_slot_registry.py index 98ddf38c661..02268dba361 100644 --- a/tests/unit/proxy/test_credential_slot_registry.py +++ b/tests/unit/proxy/test_credential_slot_registry.py @@ -124,6 +124,15 @@ DEPLOYMENT_PARAM_CLASSIFICATION: Final[Mapping[str, Classification]] = MappingPr "default_api_key_tpm_limit": NotSecret("rate limit number"), "default_api_key_rpm_limit": NotSecret("rate limit number"), "valkey_password": Unplanted(), + "anthropic_identity_token": Unplanted(), + "anthropic_identity_token_file": Unplanted(), + "anthropic_issuer_signing_key_ref": Unplanted(), + "anthropic_keycloak_token_url": NotSecret("Keycloak token endpoint URL"), + "anthropic_keycloak_client_id": NotSecret("Keycloak client identifier"), + "anthropic_keycloak_auth_method": NotSecret("name of the client authentication method"), + "anthropic_keycloak_client_secret_ref": Unplanted(), + "anthropic_keycloak_scope": NotSecret("OAuth scope string"), + "openai_identity_token_file": Unplanted(), } ) diff --git a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py index d3037e9ccfa..42a9dc441dd 100644 --- a/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py +++ b/tests/unit/proxy/test_proxy_server_endpoints_and_startup.py @@ -14859,6 +14859,57 @@ async def test_load_config_router_authorizes_fallback_targets_against_the_callin assert router.fallback_access_check is router_fallback_access_check +def test_resolve_db_litellm_param_keeps_wif_secret_pointers(monkeypatch): + from litellm.proxy.proxy_server import ProxyConfig + + monkeypatch.setenv("WIF_TEST_KC_SECRET", "kc-secret") + proxy_config = ProxyConfig() + + pointer = proxy_config._resolve_db_litellm_param( + "anthropic_keycloak_client_secret_ref", "os.environ/WIF_TEST_KC_SECRET" + ) + dereferenced = proxy_config._resolve_db_litellm_param("api_key", "os.environ/WIF_TEST_KC_SECRET") + + assert pointer == "os.environ/WIF_TEST_KC_SECRET" + assert dereferenced == "kc-secret" + + +@pytest.mark.asyncio +async def test_load_config_keeps_wif_secret_pointers_on_config_models(tmp_path, monkeypatch): + from litellm.proxy.proxy_server import ProxyConfig + + monkeypatch.setenv("WIF_TEST_SIGNING_KEY", "-----BEGIN PRIVATE KEY-----") + monkeypatch.setenv("WIF_TEST_FDRL", "fdrl_from_env") + config_file = tmp_path / "config.yaml" + config_file.write_text( + yaml.dump( + { + "model_list": [ + { + "model_name": "claude-wif", + "litellm_params": { + "model": "anthropic/claude-haiku-4-5", + "anthropic_federation_rule_id": "os.environ/WIF_TEST_FDRL", + "anthropic_identity_source": "internal_issuer", + "anthropic_issuer_url": "https://litellm.example", + "anthropic_issuer_audience": "https://api.anthropic.com", + "anthropic_issuer_signing_key_ref": "os.environ/WIF_TEST_SIGNING_KEY", + }, + } + ] + } + ) + ) + + _router, model_list, _general_settings = await ProxyConfig().load_config( + router=None, config_file_path=str(config_file) + ) + + litellm_params = model_list[0]["litellm_params"] + assert litellm_params["anthropic_federation_rule_id"] == "fdrl_from_env" + assert litellm_params["anthropic_issuer_signing_key_ref"] == "os.environ/WIF_TEST_SIGNING_KEY" + + @pytest.mark.asyncio async def test_load_config_router_budget_checks_fallback_targets_against_the_calling_key(tmp_path, monkeypatch): """A config-loaded router refuses a paid fallback target for an over-budget caller.""" diff --git a/tests/unit/router_utils/test_fallback_event_handlers.py b/tests/unit/router_utils/test_fallback_event_handlers.py index 972c04b209f..86e283a2f2b 100644 --- a/tests/unit/router_utils/test_fallback_event_handlers.py +++ b/tests/unit/router_utils/test_fallback_event_handlers.py @@ -720,9 +720,7 @@ async def test_run_async_fallback_keeps_a_request_override_distinct_from_the_bar with pytest.raises(RuntimeError, match="fallback model also failed"): await run_async_fallback( litellm_router=router, - fallback_model_group=[ - {"model": "already-attempted", "messages": [{"role": "user", "content": "shorter"}]} - ], + fallback_model_group=[{"model": "already-attempted", "messages": [{"role": "user", "content": "shorter"}]}], original_model_group="primary-model", original_exception=RuntimeError("original failed"), max_fallbacks=3, @@ -1139,6 +1137,58 @@ class TestRunAsyncFallbackTriggersCooldown: @pytest.mark.asyncio +async def test_a_stored_fallback_target_cannot_carry_a_federation_field(): + """A dict fallback target is merged into kwargs, and kwargs beat the deployment's own params, + so a stored key/team/global fallback could otherwise set the workspace a federation token is + minted for. The request itself is already forbidden to carry these, and a stored setting is + not a more trusted source than the request.""" + with pytest.raises(ValueError, match="server-owned workload identity federation parameter"): + await run_async_fallback( + litellm_router=FakeRouter(), + fallback_model_group=[{"model": "anthropic-backup", "anthropic_federation_workspace_id": "wrkspc_other"}], + original_model_group="primary-model", + original_exception=RuntimeError("upstream limited request"), + max_fallbacks=3, + fallback_depth=0, + ) + + +@pytest.mark.asyncio +async def test_a_stored_fallback_target_cannot_carry_an_openai_federation_field(): + """The OpenAI identity trio is server-owned for the same reason: a stored fallback target + naming a token file would pick which workload assertion is exchanged for the bearer.""" + with pytest.raises(ValueError, match="openai_identity_token_file"): + await run_async_fallback( + litellm_router=FakeRouter(), + fallback_model_group=[ + {"model": "openai-backup", "openai_identity_token_file": "/var/run/secrets/tokens/other"} + ], + original_model_group="primary-model", + original_exception=RuntimeError("upstream limited request"), + max_fallbacks=3, + fallback_depth=0, + ) + + +@pytest.mark.asyncio +async def test_the_refusal_is_not_swallowed_as_a_fallback_error(): + """Checked before the per-target loop on purpose: inside it, the refusal would be caught as + that target's failure and the run would quietly continue to the next one.""" + with pytest.raises(ValueError, match="anthropic_issuer_signing_key_ref"): + await run_async_fallback( + litellm_router=FakeRouter(), + fallback_model_group=[ + {"model": "anthropic-backup", "anthropic_issuer_signing_key_ref": "os.environ/ADMIN_KEY"}, + "a-perfectly-fine-model", + ], + original_model_group="primary-model", + original_exception=RuntimeError("upstream limited request"), + max_fallbacks=3, + fallback_depth=0, + include_fallback_errors=True, + ) + + async def test_run_async_fallback_stamps_fallback_info_into_metadata(): """Spend logs are built from the request metadata of the nested call, so the fallback signal has to be stamped there before recursing.""" diff --git a/tests/unit/test_anthropic_skills_transformation.py b/tests/unit/test_anthropic_skills_transformation.py index 1b917f08ca9..1d6be70dd7a 100644 --- a/tests/unit/test_anthropic_skills_transformation.py +++ b/tests/unit/test_anthropic_skills_transformation.py @@ -27,9 +27,7 @@ FAKE_API_KEY = "sk-ant-test-key-1234" FAKE_API_BASE = "https://api.anthropic.com" -def _make_mock_response( - json_data: dict, status_code: int = 200, method: str = "POST" -) -> httpx.Response: +def _make_mock_response(json_data: dict, status_code: int = 200, method: str = "POST") -> httpx.Response: return httpx.Response( status_code=status_code, json=json_data, @@ -111,9 +109,7 @@ class TestAnthropicSkillsConfigHeaderValidation: "litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key", return_value=FAKE_API_KEY, ): - headers = self.config.validate_environment( - headers={}, litellm_params=self._make_litellm_params() - ) + headers = self.config.validate_environment(headers={}, litellm_params=self._make_litellm_params()) assert headers["x-api-key"] == FAKE_API_KEY def test_sets_anthropic_version_header(self): @@ -121,9 +117,7 @@ class TestAnthropicSkillsConfigHeaderValidation: "litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key", return_value=FAKE_API_KEY, ): - headers = self.config.validate_environment( - headers={}, litellm_params=self._make_litellm_params() - ) + headers = self.config.validate_environment(headers={}, litellm_params=self._make_litellm_params()) assert headers["anthropic-version"] == "2023-06-01" def test_sets_skills_beta_header(self): @@ -131,12 +125,12 @@ class TestAnthropicSkillsConfigHeaderValidation: "litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key", return_value=FAKE_API_KEY, ): - headers = self.config.validate_environment( - headers={}, litellm_params=self._make_litellm_params() - ) + headers = self.config.validate_environment(headers={}, litellm_params=self._make_litellm_params()) assert headers["anthropic-beta"] == ANTHROPIC_SKILLS_API_BETA_VERSION - def test_merges_existing_beta_header_string(self): + def test_merges_existing_beta_header_into_string(self): + """The merged value must stay a comma-separated string: a list value makes + httpx.Headers raise TypeError when the request is built.""" with patch( "litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key", return_value=FAKE_API_KEY, @@ -145,21 +139,26 @@ class TestAnthropicSkillsConfigHeaderValidation: headers={"anthropic-beta": "other-beta-2024-01-01"}, litellm_params=self._make_litellm_params(), ) - assert isinstance(headers["anthropic-beta"], list) - assert "other-beta-2024-01-01" in headers["anthropic-beta"] - assert ANTHROPIC_SKILLS_API_BETA_VERSION in headers["anthropic-beta"] + assert isinstance(headers["anthropic-beta"], str) + betas = set(headers["anthropic-beta"].split(",")) + assert {"other-beta-2024-01-01", ANTHROPIC_SKILLS_API_BETA_VERSION} <= betas + httpx.Headers(headers) - def test_merges_existing_beta_header_list(self): + def test_oauth_key_beta_merges_without_crashing_httpx(self): + """Regression: an sk-ant-oat/WIF auth header carries its own anthropic-beta; + the old list-building merge produced a Python list that crashed httpx.""" with patch( "litellm.llms.anthropic.common_utils.AnthropicModelInfo.get_api_key", - return_value=FAKE_API_KEY, + return_value="sk-ant-oat01-fake-skills-token", ): headers = self.config.validate_environment( - headers={"anthropic-beta": ["other-beta-2024-01-01"]}, - litellm_params=self._make_litellm_params(), + headers={}, litellm_params=self._make_litellm_params(api_key=None) ) - assert ANTHROPIC_SKILLS_API_BETA_VERSION in headers["anthropic-beta"] - assert "other-beta-2024-01-01" in headers["anthropic-beta"] + assert headers["authorization"] == "Bearer sk-ant-oat01-fake-skills-token" + assert isinstance(headers["anthropic-beta"], str) + betas = set(headers["anthropic-beta"].split(",")) + assert {"oauth-2025-04-20", ANTHROPIC_SKILLS_API_BETA_VERSION} <= betas + httpx.Headers(headers) def test_does_not_duplicate_beta_header(self): with patch( @@ -170,11 +169,7 @@ class TestAnthropicSkillsConfigHeaderValidation: headers={"anthropic-beta": ANTHROPIC_SKILLS_API_BETA_VERSION}, litellm_params=self._make_litellm_params(), ) - beta = headers["anthropic-beta"] - if isinstance(beta, list): - assert beta.count(ANTHROPIC_SKILLS_API_BETA_VERSION) == 1 - else: - assert beta == ANTHROPIC_SKILLS_API_BETA_VERSION + assert headers["anthropic-beta"] == ANTHROPIC_SKILLS_API_BETA_VERSION def test_raises_without_api_key(self): with patch( @@ -182,9 +177,7 @@ class TestAnthropicSkillsConfigHeaderValidation: return_value=None, ): with pytest.raises(ValueError, match="ANTHROPIC_API_KEY"): - self.config.validate_environment( - headers={}, litellm_params=self._make_litellm_params(api_key=None) - ) + self.config.validate_environment(headers={}, litellm_params=self._make_litellm_params(api_key=None)) class TestAnthropicSkillsConfigCreateRequestTransformation: @@ -275,9 +268,7 @@ class TestAnthropicSkillsConfigResponseTransformation: def test_create_skill_response_parses_skill(self): payload = _make_skill_payload() raw = _make_mock_response(payload) - skill = self.config.transform_create_skill_response( - raw_response=raw, logging_obj=self.logging_obj - ) + skill = self.config.transform_create_skill_response(raw_response=raw, logging_obj=self.logging_obj) assert isinstance(skill, Skill) assert skill.id == "skill_abc123" assert skill.source == "custom" @@ -286,9 +277,7 @@ class TestAnthropicSkillsConfigResponseTransformation: def test_get_skill_response_parses_skill(self): payload = _make_skill_payload(id="skill_xyz", display_title="Another") raw = _make_mock_response(payload, method="GET") - skill = self.config.transform_get_skill_response( - raw_response=raw, logging_obj=self.logging_obj - ) + skill = self.config.transform_get_skill_response(raw_response=raw, logging_obj=self.logging_obj) assert isinstance(skill, Skill) assert skill.id == "skill_xyz" assert skill.display_title == "Another" @@ -300,9 +289,7 @@ class TestAnthropicSkillsConfigResponseTransformation: "next_page": None, } raw = _make_mock_response(payload, method="GET") - result = self.config.transform_list_skills_response( - raw_response=raw, logging_obj=self.logging_obj - ) + result = self.config.transform_list_skills_response(raw_response=raw, logging_obj=self.logging_obj) assert isinstance(result, ListSkillsResponse) assert len(result.data) == 2 assert result.data[0].id == "skill_abc123" @@ -316,18 +303,14 @@ class TestAnthropicSkillsConfigResponseTransformation: "next_page": "page_token_xyz", } raw = _make_mock_response(payload, method="GET") - result = self.config.transform_list_skills_response( - raw_response=raw, logging_obj=self.logging_obj - ) + result = self.config.transform_list_skills_response(raw_response=raw, logging_obj=self.logging_obj) assert result.has_more is True assert result.next_page == "page_token_xyz" def test_delete_skill_response_parses_correctly(self): payload = {"id": "skill_abc123", "type": "skill_deleted"} raw = _make_mock_response(payload, method="DELETE") - result = self.config.transform_delete_skill_response( - raw_response=raw, logging_obj=self.logging_obj - ) + result = self.config.transform_delete_skill_response(raw_response=raw, logging_obj=self.logging_obj) assert isinstance(result, DeleteSkillResponse) assert result.id == "skill_abc123" assert result.type == "skill_deleted" @@ -341,8 +324,6 @@ class TestAnthropicSkillsConfigResponseTransformation: "type": "skill", } raw = _make_mock_response(payload) - skill = self.config.transform_create_skill_response( - raw_response=raw, logging_obj=self.logging_obj - ) + skill = self.config.transform_create_skill_response(raw_response=raw, logging_obj=self.logging_obj) assert skill.display_title is None assert skill.latest_version is None diff --git a/tests/unit/test_lazy_imports.py b/tests/unit/test_lazy_imports.py index 07ead78207b..10986f8a140 100644 --- a/tests/unit/test_lazy_imports.py +++ b/tests/unit/test_lazy_imports.py @@ -39,6 +39,7 @@ from litellm._lazy_imports import ( UTILS_MODULE_NAMES, _lazy_import_utils_module, ) +from tests.test_litellm_rust.support.child_interpreter import run_child_interpreter def test_import_litellm_does_not_load_fastapi_or_bpe_table(): @@ -365,3 +366,14 @@ def test_utils_module_lazy_imports(): assert name in utils_globals _verify_only_requested_name_imported_in_utils(name, UTILS_MODULE_NAMES) + + +@pytest.mark.parametrize( + "module", + ["litellm.litellm_core_utils.get_litellm_params", "litellm.batches.batch_utils", "litellm.types.utils"], +) +def test_kwargs_funnel_and_its_importers_load_first_in_fresh_process(module: str): + """These modules are often the first to pull in litellm.types.utils, and the WIF key sets shared between + the funnel and all_litellm_params must not turn that into a cycle.""" + result = run_child_interpreter(f"import {module}", timeout=120) + assert result.returncode == 0, result.stderr diff --git a/tests/unit/test_router/test_router.py b/tests/unit/test_router/test_router.py index 5d9b37e24fc..d0115593e46 100644 --- a/tests/unit/test_router/test_router.py +++ b/tests/unit/test_router/test_router.py @@ -19147,6 +19147,65 @@ async def test_router_subclass_overriding_async_get_healthy_deployments_with_the assert response.choices[0].message.content == "hi" +def test_get_deployment_credentials_with_provider_preserves_anthropic_wif_params(): + """ + Test that get_deployment_credentials_with_provider preserves a litellm_params-configured + Anthropic workload identity federation setup (both the legacy token_file fields and the + Phase 1 internal_issuer/keycloak identity-source fields) so files/batches/passthrough + deployments using WIF do not silently fall back to a missing credential. + """ + wif_params = { + "anthropic_federation_rule_id": "fdrl_deployment", + "anthropic_organization_id": "org-deployment", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": "https://keycloak.internal.example/realms/r/protocol/openid-connect/token", + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_client_secret_ref": "oidc/env/KEYCLOAK_CLIENT_SECRET", + } + router = litellm.Router( + model_list=[ + { + "model_name": "anthropic-wif-model", + "litellm_params": { + "model": "anthropic/claude-sonnet-4-5", + **wif_params, + }, + } + ], + ) + + credentials = router.get_deployment_credentials_with_provider(model_id="anthropic-wif-model") + + assert credentials is not None + for key, value in wif_params.items(): + assert credentials.get(key) == value, key + + +def test_router_keeps_wif_secret_pointers_unresolved(monkeypatch): + monkeypatch.setenv("WIF_TEST_KC_SECRET", "kc-secret") + monkeypatch.setenv("WIF_TEST_FDRL", "fdrl_from_env") + router = Router( + model_list=[ + { + "model_name": "claude-wif", + "litellm_params": { + "model": "anthropic/claude-haiku-4-5", + "anthropic_federation_rule_id": "os.environ/WIF_TEST_FDRL", + "anthropic_identity_source": "keycloak", + "anthropic_keycloak_token_url": "https://keycloak.example/token", + "anthropic_keycloak_client_id": "litellm", + "anthropic_keycloak_client_secret_ref": "os.environ/WIF_TEST_KC_SECRET", + }, + } + ] + ) + + litellm_params = router.get_model_list()[0]["litellm_params"] + + assert litellm_params["anthropic_federation_rule_id"] == "fdrl_from_env" + assert litellm_params["anthropic_keycloak_client_secret_ref"] == "os.environ/WIF_TEST_KC_SECRET" + + @pytest.mark.asyncio async def test_failure_rpm_increment_declares_the_router_usage_key_family(): """The RPM bump a failed call still earns is router usage bookkeeping, so its Redis span diff --git a/tests/unit/types/test_litellm_params.py b/tests/unit/types/test_litellm_params.py index e3bbae39468..8d163731a51 100644 --- a/tests/unit/types/test_litellm_params.py +++ b/tests/unit/types/test_litellm_params.py @@ -95,6 +95,27 @@ CONNECTION_NAMES: Final = ( "s3_secret_access_key", "s3_encryption_key_id", "bedrock_tags", + "anthropic_federation_rule_id", + "anthropic_organization_id", + "anthropic_service_account_id", + "anthropic_federation_workspace_id", + "anthropic_identity_token_file", + "anthropic_identity_token", + "anthropic_identity_source", + "anthropic_issuer_url", + "anthropic_issuer_subject", + "anthropic_issuer_audience", + "anthropic_issuer_ttl_seconds", + "anthropic_issuer_signing_key_ref", + "anthropic_keycloak_token_url", + "anthropic_keycloak_client_id", + "anthropic_keycloak_auth_method", + "anthropic_keycloak_client_secret_ref", + "anthropic_keycloak_scope", + "anthropic_disable_workload_identity_federation", + "openai_identity_provider_id", + "openai_service_account_id", + "openai_identity_token_file", ) OPTION_NAMES: Final = ( @@ -491,6 +512,12 @@ TYPE_HINT_NAMESPACE: Final[Mapping[str, object]] = { LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { litellm_params.ProviderConnection: {"api_key": "k", "request_timeout": 1.5}, litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": "arn", "bedrock_tags": ({"k": "v"},)}, + litellm_params.AnthropicFederationConnection: { + "anthropic_federation_rule_id": "fdrl_1", + "anthropic_issuer_ttl_seconds": 300, + "anthropic_disable_workload_identity_federation": True, + }, + litellm_params.OpenAIFederationConnection: {"openai_identity_provider_id": "idp_1"}, litellm_params.DispatchOptions: {"custom_llm_provider": "openai"}, litellm_params.RoutingOptions: { "fallbacks": [{"model": "gpt-4o", "api_key": "k", "temperature": 0}], @@ -526,6 +553,8 @@ LEAF_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { LEAF_BAD_SAMPLES: Final[Mapping[type, Mapping[str, object]]] = { litellm_params.ProviderConnection: {"api_key": 1}, litellm_params.BedrockBatchConnection: {"aws_batch_role_arn": 1}, + litellm_params.AnthropicFederationConnection: {"anthropic_issuer_ttl_seconds": "300"}, + litellm_params.OpenAIFederationConnection: {"openai_identity_provider_id": 1}, litellm_params.DispatchOptions: {"custom_llm_provider": 1}, litellm_params.RoutingOptions: {"num_retries": "2"}, litellm_params.DeploymentOptions: {"rpm": "2"}, @@ -620,6 +649,24 @@ def test_routing_options_accept_every_strategy_the_router_accepts(strategy: str) NAMES_SHARED_WITH_TYPED_MODELS: Final[Mapping[str, tuple[str, ...]]] = MappingProxyType( { "credentials": ( + "anthropic_disable_workload_identity_federation", + "anthropic_federation_rule_id", + "anthropic_federation_workspace_id", + "anthropic_identity_source", + "anthropic_identity_token", + "anthropic_identity_token_file", + "anthropic_issuer_audience", + "anthropic_issuer_signing_key_ref", + "anthropic_issuer_subject", + "anthropic_issuer_ttl_seconds", + "anthropic_issuer_url", + "anthropic_keycloak_auth_method", + "anthropic_keycloak_client_id", + "anthropic_keycloak_client_secret_ref", + "anthropic_keycloak_scope", + "anthropic_keycloak_token_url", + "anthropic_organization_id", + "anthropic_service_account_id", "api_base", "api_key", "api_version", @@ -630,6 +677,9 @@ NAMES_SHARED_WITH_TYPED_MODELS: Final[Mapping[str, tuple[str, ...]]] = MappingPr "bedrock_tags", "client_id", "client_secret", + "openai_identity_provider_id", + "openai_identity_token_file", + "openai_service_account_id", "region_name", "s3_access_key_id", "s3_bucket_name", diff --git a/tests/unit/types/test_router.py b/tests/unit/types/test_router.py index 4881b094cd6..d2817ae6a90 100644 --- a/tests/unit/types/test_router.py +++ b/tests/unit/types/test_router.py @@ -5,12 +5,23 @@ from pydantic import ValidationError from litellm.types.router import ( SPECIAL_MODEL_INFO_PARAMS, + CredentialLiteLLMParams, Deployment, GenericLiteLLMParams, LiteLLM_Params, ModelInfo, + holds_secret_pointer, + reject_server_owned_wif_params, + server_owned_wif_fields_named, + server_owned_wif_fields_present, +) +from litellm.types.utils import ( + CustomPricingLiteLLMParams, + MirroredPricingParams, + anthropic_wif_litellm_params, + openai_wif_litellm_params, + server_owned_wif_litellm_params, ) -from litellm.types.utils import CustomPricingLiteLLMParams, MirroredPricingParams def test_model_info_declares_mirrored_pricing_fields(): @@ -234,3 +245,103 @@ def test_model_info_rejects_offset_aware_access_window_times(): {"start": "22:00+05:00", "end": "06:00", "timezone": "UTC", "team_ids": ["t"]} ], ) + + +def test_credential_litellm_params_declares_every_anthropic_wif_field(): + """Without these, get_deployment_credentials_with_provider round-trips litellm_params + through a strict Pydantic dump and silently drops every WIF field before files/batches/ + passthrough callers see it -- the same #30235-shaped gap azure_ad_token closed above.""" + for field in anthropic_wif_litellm_params: + assert field in CredentialLiteLLMParams.model_fields, field + + +def test_anthropic_wif_fields_round_trip_through_model_dump(): + values = {field: f"value-for-{field}" for field in anthropic_wif_litellm_params} + values["anthropic_issuer_ttl_seconds"] = 300 + values["anthropic_disable_workload_identity_federation"] = True + + dumped = CredentialLiteLLMParams(**values).model_dump(exclude_none=True) + + for field, value in values.items(): + assert dumped[field] == value, field + + +def test_server_owned_wif_fields_present_reports_only_set_fields(): + assert server_owned_wif_fields_present({}) == () + assert server_owned_wif_fields_present({"model": "gpt-4o"}) == () + assert server_owned_wif_fields_present( + {"anthropic_keycloak_token_url": "https://idp.example/token", "model": "gpt-4o"} + ) == ("anthropic_keycloak_token_url",) + + +def test_server_owned_wif_fields_present_is_derived_from_the_shared_list(): + """A non-admin persistence gate built on this must automatically cover a field added + later to server_owned_wif_litellm_params, not just the fields known when the gate was + written -- so this must read the shared list rather than a hand-copied one.""" + values = {field: "set" for field in server_owned_wif_litellm_params} + assert set(server_owned_wif_fields_present(values)) == set(server_owned_wif_litellm_params) + + +def test_server_owned_wif_fields_named_reports_keys_whatever_their_value(): + """The credential write gates must see a key a caller sets to ``None``: the federation + resolver reacts to the key's presence, not its value, so ``{"anthropic_issuer_url": None}`` + wedges every deployment referencing the credential once persisted.""" + assert server_owned_wif_fields_named({}) == () + assert server_owned_wif_fields_named({"model": "gpt-4o"}) == () + assert server_owned_wif_fields_named({"anthropic_issuer_url": None}) == ("anthropic_issuer_url",) + assert server_owned_wif_fields_present({"anthropic_issuer_url": None}) == () + assert server_owned_wif_fields_named(("anthropic_keycloak_token_url", "api_key")) == ( + "anthropic_keycloak_token_url", + ) + + +def test_server_owned_wif_fields_named_is_derived_from_the_shared_list(): + assert set(server_owned_wif_fields_named(frozenset(server_owned_wif_litellm_params))) == set( + server_owned_wif_litellm_params + ) + + +@pytest.mark.parametrize("param_name", ["anthropic_issuer_signing_key_ref", "anthropic_keycloak_client_secret_ref"]) +def test_wif_ref_fields_hold_secret_pointers(param_name: str): + assert holds_secret_pointer(param_name) + + +@pytest.mark.parametrize("param_name", ["api_key", "anthropic_federation_rule_id", "anthropic_identity_token"]) +def test_dereferenced_fields_do_not_hold_secret_pointers(param_name: str): + assert not holds_secret_pointer(param_name) + + +def test_credential_litellm_params_declares_every_openai_wif_field(): + for field in openai_wif_litellm_params: + assert field in CredentialLiteLLMParams.model_fields, field + + +def test_openai_wif_fields_round_trip_through_model_dump(): + values = {field: f"value-for-{field}" for field in openai_wif_litellm_params} + + dumped = CredentialLiteLLMParams(**values).model_dump(exclude_none=True) + + for field, value in values.items(): + assert dumped[field] == value, field + + +def test_server_owned_registry_is_anthropic_plus_openai(): + assert server_owned_wif_litellm_params == anthropic_wif_litellm_params + openai_wif_litellm_params + assert set(openai_wif_litellm_params) == { + "openai_identity_provider_id", + "openai_service_account_id", + "openai_identity_token_file", + } + + +def test_server_owned_wif_fields_present_reports_openai_fields(): + assert server_owned_wif_fields_present( + {"openai_identity_token_file": "/var/run/secrets/tokens/openai", "model": "gpt-4o"} + ) == ("openai_identity_token_file",) + assert server_owned_wif_fields_named({"openai_service_account_id": None}) == ("openai_service_account_id",) + + +@pytest.mark.parametrize("param_name", openai_wif_litellm_params) +def test_reject_server_owned_wif_params_names_each_openai_field(param_name: str): + with pytest.raises(ValueError, match=param_name): + reject_server_owned_wif_params({param_name: "client-supplied"}) diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index d1ef1e92008..62c5990c348 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -3808,6 +3808,28 @@ export interface paths { patch: operations["update_credential_credentials__credential_name__patch"]; trace?: never; }; + "/credentials/{credential_name}/jwks": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * Get Credential Internal Issuer Jwks + * @description Export the public JWKS for an anthropic ``internal_issuer`` credential, so the operator can + * register it on the Anthropic federation issuer from the UI. Never touches the private signing + * key: only its derived public JWKS leaves this process. 404s for any other credential shape. + */ + get: operations["get_credential_internal_issuer_jwks_credentials__credential_name__jwks_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/cursor/chat/completions": { parameters: { query?: never; @@ -31454,6 +31476,8 @@ export interface components { }; /** Deployment */ Deployment: { + /** Blocked */ + blocked?: boolean | null; litellm_params: components["schemas"]["LiteLLM_Params"]; model_info: components["schemas"]["litellm__types__router__ModelInfo"]; /** Model Name */ @@ -35010,6 +35034,42 @@ export interface components { annotation_cost_per_page?: number | null; /** Annotation Cost Per Page Batches */ annotation_cost_per_page_batches?: number | null; + /** Anthropic Disable Workload Identity Federation */ + anthropic_disable_workload_identity_federation?: boolean | null; + /** Anthropic Federation Rule Id */ + anthropic_federation_rule_id?: string | null; + /** Anthropic Federation Workspace Id */ + anthropic_federation_workspace_id?: string | null; + /** Anthropic Identity Source */ + anthropic_identity_source?: string | null; + /** Anthropic Identity Token */ + anthropic_identity_token?: string | null; + /** Anthropic Identity Token File */ + anthropic_identity_token_file?: string | null; + /** Anthropic Issuer Audience */ + anthropic_issuer_audience?: string | null; + /** Anthropic Issuer Signing Key Ref */ + anthropic_issuer_signing_key_ref?: string | null; + /** Anthropic Issuer Subject */ + anthropic_issuer_subject?: string | null; + /** Anthropic Issuer Ttl Seconds */ + anthropic_issuer_ttl_seconds?: number | null; + /** Anthropic Issuer Url */ + anthropic_issuer_url?: string | null; + /** Anthropic Keycloak Auth Method */ + anthropic_keycloak_auth_method?: string | null; + /** Anthropic Keycloak Client Id */ + anthropic_keycloak_client_id?: string | null; + /** Anthropic Keycloak Client Secret Ref */ + anthropic_keycloak_client_secret_ref?: string | null; + /** Anthropic Keycloak Scope */ + anthropic_keycloak_scope?: string | null; + /** Anthropic Keycloak Token Url */ + anthropic_keycloak_token_url?: string | null; + /** Anthropic Organization Id */ + anthropic_organization_id?: string | null; + /** Anthropic Service Account Id */ + anthropic_service_account_id?: string | null; /** Api Base */ api_base?: string | null; /** Api Key */ @@ -35273,6 +35333,12 @@ export interface components { ocr_cost_per_page?: number | null; /** Ocr Cost Per Page Batches */ ocr_cost_per_page_batches?: number | null; + /** Openai Identity Provider Id */ + openai_identity_provider_id?: string | null; + /** Openai Identity Token File */ + openai_identity_token_file?: string | null; + /** Openai Service Account Id */ + openai_service_account_id?: string | null; /** Organization */ organization?: string | null; /** Otpm */ @@ -48123,6 +48189,8 @@ export interface components { credential_values?: { [key: string]: unknown; } | null; + /** Credential Values To Delete */ + credential_values_to_delete?: string[] | null; /** Model Id */ model_id?: string | null; }; @@ -50386,6 +50454,42 @@ export interface components { annotation_cost_per_page?: number | null; /** Annotation Cost Per Page Batches */ annotation_cost_per_page_batches?: number | null; + /** Anthropic Disable Workload Identity Federation */ + anthropic_disable_workload_identity_federation?: boolean | null; + /** Anthropic Federation Rule Id */ + anthropic_federation_rule_id?: string | null; + /** Anthropic Federation Workspace Id */ + anthropic_federation_workspace_id?: string | null; + /** Anthropic Identity Source */ + anthropic_identity_source?: string | null; + /** Anthropic Identity Token */ + anthropic_identity_token?: string | null; + /** Anthropic Identity Token File */ + anthropic_identity_token_file?: string | null; + /** Anthropic Issuer Audience */ + anthropic_issuer_audience?: string | null; + /** Anthropic Issuer Signing Key Ref */ + anthropic_issuer_signing_key_ref?: string | null; + /** Anthropic Issuer Subject */ + anthropic_issuer_subject?: string | null; + /** Anthropic Issuer Ttl Seconds */ + anthropic_issuer_ttl_seconds?: number | null; + /** Anthropic Issuer Url */ + anthropic_issuer_url?: string | null; + /** Anthropic Keycloak Auth Method */ + anthropic_keycloak_auth_method?: string | null; + /** Anthropic Keycloak Client Id */ + anthropic_keycloak_client_id?: string | null; + /** Anthropic Keycloak Client Secret Ref */ + anthropic_keycloak_client_secret_ref?: string | null; + /** Anthropic Keycloak Scope */ + anthropic_keycloak_scope?: string | null; + /** Anthropic Keycloak Token Url */ + anthropic_keycloak_token_url?: string | null; + /** Anthropic Organization Id */ + anthropic_organization_id?: string | null; + /** Anthropic Service Account Id */ + anthropic_service_account_id?: string | null; /** Api Base */ api_base?: string | null; /** Api Key */ @@ -50649,6 +50753,12 @@ export interface components { ocr_cost_per_page?: number | null; /** Ocr Cost Per Page Batches */ ocr_cost_per_page_batches?: number | null; + /** Openai Identity Provider Id */ + openai_identity_provider_id?: string | null; + /** Openai Identity Token File */ + openai_identity_token_file?: string | null; + /** Openai Service Account Id */ + openai_service_account_id?: string | null; /** Organization */ organization?: string | null; /** Otpm */ @@ -56682,6 +56792,38 @@ export interface operations { }; }; }; + get_credential_internal_issuer_jwks_credentials__credential_name__jwks_get: { + parameters: { + query?: never; + header?: never; + path: { + /** @description The credential name, percent-decoded; may contain slashes */ + credential_name: string; + }; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": unknown; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; cursor_chat_completions_cursor_chat_completions_post: { parameters: { query?: never; From f850b2c324226759f2abae2056e1ad25dac56c22 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Sun, 4 Oct 2026 00:09:58 +0000 Subject: [PATCH 12/18] test(integration): exact four-part translation cases on a shared fake provider and shared YAML deployment (#44451) * test(integration): exact four-part translation cases on a shared fake provider and shared YAML deployment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): compare every non-transport provider header and check for late provider requests at session end Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): name TranslationTestCase fields after litellm and provider sides and drop regressions Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): prefix checked TranslationTestCase fields with expected_ and name the fake reply mock_provider_response Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(integration): name TranslationTestCase fields in the translation README Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): add a claude-opus-5-5 base case and deployment next to claude-sonnet-4-6 Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): name translation cases _TEST_CASE and document the naming rule Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * docs(integration): move translation test rules into tests/integration/translation Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/README.md | 2 + tests/integration/_support/manifest.py | 1 + tests/integration/_support/provider.py | 44 ++++++ tests/integration/conftest.py | 28 ++++ tests/integration/proxy_config.yaml | 12 +- tests/integration/run.py | 2 +- tests/integration/translation/AGENTS.md | 21 +++ tests/integration/translation/README.md | 7 + tests/integration/translation/__init__.py | 0 tests/integration/translation/case.py | 24 +++ tests/integration/translation/conftest.py | 3 + .../translation/messages/__init__.py | 0 .../translation/messages/bases/__init__.py | 0 .../translation/messages/bases/anthropic.py | 139 ++++++++++++++++++ .../translation/messages/basic/__init__.py | 0 .../basic/test_messages_basic_anthropic.py | 11 ++ .../messages/reasoning/__init__.py | 0 .../test_messages_reasoning_anthropic.py | 72 +++++++++ tests/integration/translation/runner.py | 21 +++ 19 files changed, 385 insertions(+), 2 deletions(-) create mode 100644 tests/integration/_support/provider.py create mode 100644 tests/integration/translation/AGENTS.md create mode 100644 tests/integration/translation/README.md create mode 100644 tests/integration/translation/__init__.py create mode 100644 tests/integration/translation/case.py create mode 100644 tests/integration/translation/conftest.py create mode 100644 tests/integration/translation/messages/__init__.py create mode 100644 tests/integration/translation/messages/bases/__init__.py create mode 100644 tests/integration/translation/messages/bases/anthropic.py create mode 100644 tests/integration/translation/messages/basic/__init__.py create mode 100644 tests/integration/translation/messages/basic/test_messages_basic_anthropic.py create mode 100644 tests/integration/translation/messages/reasoning/__init__.py create mode 100644 tests/integration/translation/messages/reasoning/test_messages_reasoning_anthropic.py create mode 100644 tests/integration/translation/runner.py diff --git a/tests/integration/README.md b/tests/integration/README.md index ef418c50759..fe31e215f8a 100644 --- a/tests/integration/README.md +++ b/tests/integration/README.md @@ -20,6 +20,8 @@ There is no per-node manifest. A positional argument is a file of the group or a Provider sentinels currently use the controlled server, not live recordings. The provider shard also runs the existing strict replay controls for changed requests, exhausted interactions, leftover interactions and no provider connection. Future recorded scenarios must use that replay-only implementation; missing recordings cannot fall back to a real provider. The observation endpoint is destructive and the current selection runs serially against one owned upstream +Translation tests in `translation/` compare exact provider requests and LiteLLM responses against a shared fake provider; `translation/README.md` has their rules + Fixtures must contain synthetic data only. Keep private incident records and source documents out of code, fixtures, logs and PR descriptions Database cases own their temporary schemas, roles, constraints and proxy processes. They prove reader-versus-writer execution with PostgreSQL lock observations, exercise real transaction wait limits and verify rollback after a reached database failure diff --git a/tests/integration/_support/manifest.py b/tests/integration/_support/manifest.py index a4a86a21568..fb17ebb3864 100644 --- a/tests/integration/_support/manifest.py +++ b/tests/integration/_support/manifest.py @@ -11,6 +11,7 @@ OWNED_DIRECTORIES: Final = frozenset( "providers", "streaming", "messages_endpoint", + "translation", "configuration", "mcp", "observability", diff --git a/tests/integration/_support/provider.py b/tests/integration/_support/provider.py new file mode 100644 index 00000000000..278cd054c61 --- /dev/null +++ b/tests/integration/_support/provider.py @@ -0,0 +1,44 @@ +"""The fake provider shared by every integration test that takes the `provider` fixture. + +Deployments in `proxy_config.yaml` point at `PROVIDER_URL`, so one server answers for all of them. A test +queues the replies it expects with `expect` and reads what the proxy sent with `received`. Tests run one at a +time against it; the `provider` fixture checks nothing is left over between tests. +""" + +from __future__ import annotations + +from collections import deque +from collections.abc import Iterator +from contextlib import contextmanager +from dataclasses import dataclass, field +from typing import Final + +from tests.integration._support.wire import Reply, Request, Wire, wire_server + +PROVIDER_PORT: Final = 8191 +PROVIDER_URL: Final = f"http://127.0.0.1:{PROVIDER_PORT}" +_UNQUEUED: Final = Reply(status=500, body=b'{"error": "the shared fake provider has no reply queued for this request"}') + + +@dataclass(slots=True) +class SharedProvider: + wire: Wire + replies: deque[Reply] + last_test: str | None = field(default=None) + + def expect(self, *replies: Reply) -> None: + self.replies.extend(replies) + + def received(self) -> tuple[Request, ...]: + return self.wire.drain() + + +@contextmanager +def shared_provider() -> Iterator[SharedProvider]: + replies: Final[deque[Reply]] = deque() + + def respond(request: Request) -> Reply: + return replies.popleft() if replies else _UNQUEUED + + with wire_server(respond, port=PROVIDER_PORT) as wire: + yield SharedProvider(wire, replies) diff --git a/tests/integration/conftest.py b/tests/integration/conftest.py index 14e31441a95..116bc237018 100644 --- a/tests/integration/conftest.py +++ b/tests/integration/conftest.py @@ -15,6 +15,7 @@ from redis import Redis from tests.integration._support.client import Gateway, eventually, gateway_from_environment from tests.integration._support.generation import LIFECYCLE_SETTINGS from tests.integration._support.manifest import OWNED_DIRECTORIES +from tests.integration._support.provider import SharedProvider, shared_provider from tests.integration._support.routing import RoutingPlugin from tests.integration.run import GITHUB_FILES @@ -130,6 +131,33 @@ def gateway() -> Iterator[Gateway]: yield value +@pytest.fixture(scope="session") +def shared_provider_server() -> Iterator[SharedProvider]: + if os.environ.get("PYTEST_XDIST_WORKER"): + pytest.fail("the shared fake provider needs tests to run one at a time; this group runs under pytest-xdist") + with shared_provider() as server: + yield server + late: Final = server.received() + assert late == (), f"the shared fake provider got {[item.target for item in late]} after {server.last_test} finished" + + +@pytest.fixture +def provider(shared_provider_server: SharedProvider, request: pytest.FixtureRequest) -> Iterator[SharedProvider]: + stray: Final = shared_provider_server.received() + shared_provider_server.replies.clear() + assert stray == (), ( + f"the shared fake provider got {[item.target for item in stray]} " + f"after {shared_provider_server.last_test} finished" + ) + yield shared_provider_server + shared_provider_server.last_test = request.node.nodeid + unused: Final = len(shared_provider_server.replies) + unread: Final = shared_provider_server.received() + shared_provider_server.replies.clear() + assert unused == 0, f"{unused} queued provider replies were never requested" + assert unread == (), f"the test never read the provider requests {[item.target for item in unread]}" + + @pytest.fixture def peer(gateway: Gateway) -> Iterator[Gateway]: url: Final = os.environ["INTEGRATION_PEER_URL"] diff --git a/tests/integration/proxy_config.yaml b/tests/integration/proxy_config.yaml index 085adea81ac..4c9db394565 100644 --- a/tests/integration/proxy_config.yaml +++ b/tests/integration/proxy_config.yaml @@ -1,4 +1,14 @@ -model_list: [] +model_list: + - model_name: anthropic/claude-opus-5-5 + litellm_params: + model: anthropic/claude-opus-5-5 + api_base: http://127.0.0.1:8191 + api_key: synthetic-anthropic-key + - model_name: anthropic/claude-sonnet-4-6 + litellm_params: + model: anthropic/claude-sonnet-4-6 + api_base: http://127.0.0.1:8191 + api_key: synthetic-anthropic-key general_settings: master_key: os.environ/LITELLM_MASTER_KEY database_url: os.environ/DATABASE_URL diff --git a/tests/integration/run.py b/tests/integration/run.py index aca21ec522c..df408cfc809 100644 --- a/tests/integration/run.py +++ b/tests/integration/run.py @@ -15,7 +15,7 @@ GROUPS: Final = MappingProxyType( "management": ("management", "authorization", "configuration"), "accounting": ("pricing", "spend"), "database": ("database",), - "providers": ("providers", "routing", "streaming", "messages_endpoint"), + "providers": ("providers", "routing", "streaming", "messages_endpoint", "translation"), "extensions": ("observability", "compatibility"), "mcp": ("mcp",), "sdk": ("sdk",), diff --git a/tests/integration/translation/AGENTS.md b/tests/integration/translation/AGENTS.md new file mode 100644 index 00000000000..33d176256e3 --- /dev/null +++ b/tests/integration/translation/AGENTS.md @@ -0,0 +1,21 @@ +# tests/integration/translation + +Exact translation cases on the shared fake provider. `README.md` here has the case fields, deployments and capture steps + +## Naming + +Each model gets one complete base `TranslationTestCase` in `translation//bases/.py`, +named `_TEST_CASE` after its deployment (`anthropic/claude-sonnet-4-6` is +`CLAUDE_SONNET_4_6_TEST_CASE`). A feature case is `__TEST_CASE`. Import a base under its +own name, never aliased to `BASE`, so every case shows which model it derives from + +```python +from integration.translation.messages.bases.anthropic import CLAUDE_SONNET_4_6_TEST_CASE + +CLAUDE_SONNET_4_6_THINKING_BUDGET_TEST_CASE: Final = replace( + CLAUDE_SONNET_4_6_TEST_CASE, + scenario="thinking_budget", + litellm_request={**CLAUDE_SONNET_4_6_TEST_CASE.litellm_request, "max_tokens": 2048, "thinking": ...}, + ... +) +``` diff --git a/tests/integration/translation/README.md b/tests/integration/translation/README.md new file mode 100644 index 00000000000..65b8e065926 --- /dev/null +++ b/tests/integration/translation/README.md @@ -0,0 +1,7 @@ +# Translation tests + +These tests check one request through the proxy as literals on a `TranslationTestCase`: the `litellm_endpoint` and `litellm_request` the test sends, the `expected_provider_endpoint`, `expected_provider_headers` and `expected_provider_request` the fake provider must receive, the `mock_provider_response` it answers with, and the `expected_litellm_status_code` and `expected_litellm_response` the test must get back. The runner compares the provider request body, every provider header other than transport headers, and the LiteLLM response body in full, so an added, removed, renamed or moved field fails. Folders follow the client endpoint and then the feature, for example `translation/messages/reasoning/`. Each endpoint keeps one complete base case per model in `/bases/.py`, named `_TEST_CASE` (for example `CLAUDE_SONNET_4_6_TEST_CASE`), which `/basic/` runs on its own. A feature case is `__TEST_CASE = dataclasses.replace(_TEST_CASE, ...)`, imports the base under its own name rather than as `BASE`, and lists only the fields it changes + +Deployments used by translation tests are shared by the whole suite and declared in `proxy_config.yaml` under `model_list`, with `model_name` equal to the litellm model string, `api_base: http://127.0.0.1:8191` and a synthetic key. A case names the deployment literally in its client request. The fake provider behind them is the `provider` fixture: one `wire_server` on port 8191 inside the pytest process, started the first time a test asks for it. A test queues its replies with `provider.expect(...)` and reads what the proxy sent with `provider.received()`. Tests that use it run one at a time. After each test the fixture fails if a queued reply was never requested or a received request was never read, and before each test it fails if a request arrived in between, naming the previous test. It refuses to start under pytest-xdist, so the `mcp` and `cost` groups cannot use it. Client requests carry `"cache": {"no-cache": True}` because the proxy caches responses in Redis + +Provider responses in translation cases are captured once from the real provider and stored verbatim. First run the new case against the fake provider; the body the proxy sends is the case's `expected_provider_request`. Send that exact body to the real provider endpoint with a key from the 1Password `Shared` vault (`/qa-keys`), and only store the case when the provider answers 2xx. Keep only the response body and drop every response header, since headers carry account identifiers such as the organization id and rate limits. Never print or save the request headers you sent. Before committing, check the body contains no key and no account identifier, such as an organization id, an AWS account id inside an ARN, a GCP project id or an Azure resource name, and that the prompt is synthetic. Paste the body as the case's `mock_provider_response` without shortening ids, token counts or signatures, and move long opaque values such as thinking signatures into module-level constants used by both `mock_provider_response` and `expected_litellm_response`. Do not commit the script used for the capture diff --git a/tests/integration/translation/__init__.py b/tests/integration/translation/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/translation/case.py b/tests/integration/translation/case.py new file mode 100644 index 00000000000..ac8f43439ce --- /dev/null +++ b/tests/integration/translation/case.py @@ -0,0 +1,24 @@ +from collections.abc import Mapping +from dataclasses import dataclass + +from pydantic import JsonValue + + +@dataclass(frozen=True, slots=True, kw_only=True) +class TranslationTestCase: + """One request through the proxy to a deployment in `proxy_config.yaml`: what the test sends to LiteLLM, the + exact request the provider must receive, the fake provider's reply, and the exact response LiteLLM must return.""" + + scenario: str + litellm_endpoint: str + litellm_request: Mapping[str, JsonValue] + expected_provider_endpoint: str + expected_provider_headers: Mapping[str, str] + expected_provider_request: Mapping[str, JsonValue] + mock_provider_response: Mapping[str, JsonValue] + expected_litellm_status_code: int = 200 + expected_litellm_response: Mapping[str, JsonValue] + + @property + def id(self) -> str: + return f"{self.litellm_request['model']}-{self.scenario}" diff --git a/tests/integration/translation/conftest.py b/tests/integration/translation/conftest.py new file mode 100644 index 00000000000..c0dc4462cae --- /dev/null +++ b/tests/integration/translation/conftest.py @@ -0,0 +1,3 @@ +import pytest + +pytest.register_assert_rewrite("integration.translation.runner") diff --git a/tests/integration/translation/messages/__init__.py b/tests/integration/translation/messages/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/translation/messages/bases/__init__.py b/tests/integration/translation/messages/bases/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/translation/messages/bases/anthropic.py b/tests/integration/translation/messages/bases/anthropic.py new file mode 100644 index 00000000000..0642175a739 --- /dev/null +++ b/tests/integration/translation/messages/bases/anthropic.py @@ -0,0 +1,139 @@ +from typing import Final + +from integration.translation.case import TranslationTestCase + +CLAUDE_OPUS_5_5_TEST_CASE: Final = TranslationTestCase( + scenario="basic", + litellm_endpoint="/v1/messages", + litellm_request={ + "model": "anthropic/claude-opus-5-5", + "max_tokens": 64, + "system": "You are a terse assistant.", + "messages": [{"role": "user", "content": "Say hello."}], + "cache": {"no-cache": True}, + }, + expected_provider_endpoint="/v1/messages", + expected_provider_headers={ + "x-api-key": "synthetic-anthropic-key", + "anthropic-version": "2023-06-01", + "content-type": "application/json", + }, + expected_provider_request={ + "model": "claude-opus-5-5", + "max_tokens": 64, + "stream": False, + "system": "You are a terse assistant.", + "messages": [{"role": "user", "content": "Say hello."}], + }, + mock_provider_response={ + "model": "claude-opus-5-5", + "id": "msg_011CfgC5HKvyve78CTFAw97f", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello."}], + "container": None, + "stop_reason": "end_turn", + "stop_sequence": None, + "stop_details": None, + "usage": { + "input_tokens": 23, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + "output_tokens": 6, + "output_tokens_details": {"thinking_tokens": 0}, + "service_tier": "standard", + "inference_geo": "global", + }, + "diagnostics": None, + }, + expected_litellm_response={ + "model": "anthropic/claude-opus-5-5", + "id": "msg_011CfgC5HKvyve78CTFAw97f", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello."}], + "container": None, + "stop_reason": "end_turn", + "stop_sequence": None, + "stop_details": None, + "usage": { + "input_tokens": 23, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + "output_tokens": 6, + "output_tokens_details": {"thinking_tokens": 0}, + "service_tier": "standard", + "inference_geo": "global", + }, + "diagnostics": None, + }, +) + +CLAUDE_SONNET_4_6_TEST_CASE: Final = TranslationTestCase( + scenario="basic", + litellm_endpoint="/v1/messages", + litellm_request={ + "model": "anthropic/claude-sonnet-4-6", + "max_tokens": 64, + "system": "You are a terse assistant.", + "messages": [{"role": "user", "content": "Say hello."}], + "cache": {"no-cache": True}, + }, + expected_provider_endpoint="/v1/messages", + expected_provider_headers={ + "x-api-key": "synthetic-anthropic-key", + "anthropic-version": "2023-06-01", + "content-type": "application/json", + }, + expected_provider_request={ + "model": "claude-sonnet-4-6", + "max_tokens": 64, + "stream": False, + "system": "You are a terse assistant.", + "messages": [{"role": "user", "content": "Say hello."}], + }, + mock_provider_response={ + "model": "claude-sonnet-4-6", + "id": "msg_011CffzUNHaEfzVxCh5hskBG", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello!"}], + "container": None, + "stop_reason": "end_turn", + "stop_sequence": None, + "stop_details": None, + "usage": { + "input_tokens": 18, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + "output_tokens": 5, + "service_tier": "standard", + "inference_geo": "global", + }, + "diagnostics": None, + }, + expected_litellm_response={ + "model": "anthropic/claude-sonnet-4-6", + "id": "msg_011CffzUNHaEfzVxCh5hskBG", + "type": "message", + "role": "assistant", + "content": [{"type": "text", "text": "Hello!"}], + "container": None, + "stop_reason": "end_turn", + "stop_sequence": None, + "stop_details": None, + "usage": { + "input_tokens": 18, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + "output_tokens": 5, + "service_tier": "standard", + "inference_geo": "global", + }, + "diagnostics": None, + }, +) diff --git a/tests/integration/translation/messages/basic/__init__.py b/tests/integration/translation/messages/basic/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/translation/messages/basic/test_messages_basic_anthropic.py b/tests/integration/translation/messages/basic/test_messages_basic_anthropic.py new file mode 100644 index 00000000000..3fbca2addd5 --- /dev/null +++ b/tests/integration/translation/messages/basic/test_messages_basic_anthropic.py @@ -0,0 +1,11 @@ +import pytest +from integration._support.client import Gateway +from integration._support.provider import SharedProvider +from integration.translation.case import TranslationTestCase +from integration.translation.messages.bases.anthropic import CLAUDE_OPUS_5_5_TEST_CASE, CLAUDE_SONNET_4_6_TEST_CASE +from integration.translation.runner import run + + +@pytest.mark.parametrize("case", [CLAUDE_OPUS_5_5_TEST_CASE, CLAUDE_SONNET_4_6_TEST_CASE], ids=lambda case: case.id) +def test_messages_basic_anthropic(case: TranslationTestCase, gateway: Gateway, provider: SharedProvider) -> None: + run(case, gateway, provider) diff --git a/tests/integration/translation/messages/reasoning/__init__.py b/tests/integration/translation/messages/reasoning/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/tests/integration/translation/messages/reasoning/test_messages_reasoning_anthropic.py b/tests/integration/translation/messages/reasoning/test_messages_reasoning_anthropic.py new file mode 100644 index 00000000000..555e55cc4a5 --- /dev/null +++ b/tests/integration/translation/messages/reasoning/test_messages_reasoning_anthropic.py @@ -0,0 +1,72 @@ +from dataclasses import replace +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.provider import SharedProvider +from integration.translation.case import TranslationTestCase +from integration.translation.messages.bases.anthropic import CLAUDE_SONNET_4_6_TEST_CASE +from integration.translation.runner import run + +SIGNATURE_1: Final = ( + "EpECCqgBCBIYAipAivUPApu85FYYe3+cXal8EiJOza7QGqKyekC8vDSn4oyeqGa2CrarO4abiuG7dzBXjmYR8+daw4h50ZjKmak7czIRY2xh" + "dWRlLXNvbm5ldC00LTY4AEIIdGhpbmtpbmdaJGQwMDgxZjJiLWQ5NjEtNGFhYi05ZTRjLTcxYmU3ZTA0ZTY3MJoBEwoRY2xhdWRlLXNvbm5l" + "dC00LTaoAY3fhdYGEgwJHsjNkCTlV9k1jWQaDOmeP/z67YtLTojSqCIwhRXGrNzSuGfMD1HqA72lctQCy83Wkr0u8W5lBXXn+MD6WfJGTJqM" + "1FW7qRmOMOKJKhbheeMpsTs7XvdvsiDQqgM4PAJt4cwgGAE=" +) + +CLAUDE_SONNET_4_6_THINKING_BUDGET_TEST_CASE: Final = replace( + CLAUDE_SONNET_4_6_TEST_CASE, + scenario="thinking_budget", + litellm_request={ + **CLAUDE_SONNET_4_6_TEST_CASE.litellm_request, + "max_tokens": 2048, + "thinking": {"type": "enabled", "budget_tokens": 1024}, + }, + expected_provider_request={ + **CLAUDE_SONNET_4_6_TEST_CASE.expected_provider_request, + "max_tokens": 2048, + "thinking": {"type": "enabled", "budget_tokens": 1024}, + }, + mock_provider_response={ + **CLAUDE_SONNET_4_6_TEST_CASE.mock_provider_response, + "id": "msg_011CffzUREgTzMm1dXRqP2LR", + "content": [ + {"type": "thinking", "thinking": "Hello!", "signature": SIGNATURE_1}, + {"type": "text", "text": "Hello!"}, + ], + "usage": { + "input_tokens": 47, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + "output_tokens": 15, + "output_tokens_details": {"thinking_tokens": 7}, + "service_tier": "standard", + "inference_geo": "global", + }, + }, + expected_litellm_response={ + **CLAUDE_SONNET_4_6_TEST_CASE.expected_litellm_response, + "id": "msg_011CffzUREgTzMm1dXRqP2LR", + "content": [ + {"type": "thinking", "thinking": "Hello!", "signature": SIGNATURE_1}, + {"type": "text", "text": "Hello!"}, + ], + "usage": { + "input_tokens": 47, + "cache_creation_input_tokens": 0, + "cache_read_input_tokens": 0, + "cache_creation": {"ephemeral_5m_input_tokens": 0, "ephemeral_1h_input_tokens": 0}, + "output_tokens": 15, + "output_tokens_details": {"thinking_tokens": 7}, + "service_tier": "standard", + "inference_geo": "global", + }, + }, +) + + +@pytest.mark.parametrize("case", [CLAUDE_SONNET_4_6_THINKING_BUDGET_TEST_CASE], ids=lambda case: case.id) +def test_messages_reasoning_anthropic(case: TranslationTestCase, gateway: Gateway, provider: SharedProvider) -> None: + run(case, gateway, provider) diff --git a/tests/integration/translation/runner.py b/tests/integration/translation/runner.py new file mode 100644 index 00000000000..81c858ad303 --- /dev/null +++ b/tests/integration/translation/runner.py @@ -0,0 +1,21 @@ +import json +from typing import Final + +from integration._support.client import Gateway +from integration._support.provider import SharedProvider +from integration._support.wire import Reply +from integration.translation.case import TranslationTestCase + +TRANSPORT_HEADERS: Final = frozenset({"host", "accept", "accept-encoding", "connection", "content-length", "user-agent"}) + + +def run(case: TranslationTestCase, gateway: Gateway, provider: SharedProvider) -> None: + provider.expect(Reply(body=json.dumps(case.mock_provider_response).encode())) + response: Final = gateway.request("POST", case.litellm_endpoint, case.litellm_request) + received: Final = provider.received() + assert [(request.method, request.target) for request in received] == [("POST", case.expected_provider_endpoint)] + sent: Final = received[0] + assert {name: value for name, value in sent.headers.items() if name not in TRANSPORT_HEADERS} == dict(case.expected_provider_headers) + assert json.loads(sent.body) == case.expected_provider_request + assert response.status_code == case.expected_litellm_status_code, response.text + assert response.json() == case.expected_litellm_response From 4b67a2b84571bb103e60bedd1e115ea6b2cd1aef Mon Sep 17 00:00:00 2001 From: moe-berri Date: Sat, 3 Oct 2026 17:24:18 -0700 Subject: [PATCH 13/18] feat(roi): default people and branch lists to matched accounts (#44465) * feat(roi): show matched people by default in contributor lists * fix(roi): keep matched filter tabs readable on narrow screens * fix(roi): retain spend-only users and support older browsers --- litellm/proxy/roi_calculator/README.md | 2 + .../_components/MatchedPeopleToggle.tsx | 12 ++ .../_components/ObservedDetails.tsx | 21 ++- .../ObservedROIView.integration.test.tsx | 76 +++++++++++ .../_components/ObservedReport.tsx | 68 ++++++---- .../ROICalculatorView.integration.test.tsx | 62 +++++++++ .../_components/ROICalculatorView.tsx | 14 +- .../_components/ROICalculatorViews.tsx | 33 ++--- .../_components/observedData.test.ts | 126 +++++++++++++++++- .../_components/observedData.ts | 94 ++++++++++++- .../_components/roiCalculatorData.test.ts | 55 +++++++- .../_components/roiCalculatorData.ts | 12 +- 12 files changed, 519 insertions(+), 56 deletions(-) create mode 100644 ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/MatchedPeopleToggle.tsx diff --git a/litellm/proxy/roi_calculator/README.md b/litellm/proxy/roi_calculator/README.md index 4db62fdad2b..82dd5463398 100644 --- a/litellm/proxy/roi_calculator/README.md +++ b/litellm/proxy/roi_calculator/README.md @@ -30,6 +30,8 @@ The gateway encrypts access and refresh tokens using its configured encryption k Use **Link accounts** to associate several current or historical usernames with one internal email. Each connection has a separate username field, so a GitHub username never matches a GitLab user implicitly. Saving immediately recalculates the report without fetching repositories again. Public profile emails match automatically when they resolve unambiguously to an internal user +**Matched people only** is on by default for people, merged changes, and branch lists. Turn it off to include outside contributors and their branches. Matching depends on the linked internal account, even when no spend was recorded. This switch filters the lists; summary metrics and quality signals still cover all selected repositories + Agent-authored changes count for a person only when the supported agent metadata explicitly names a requester. Repository issue counts and revert titles are quality signals, not an individual defect score Bug and regression counts combine repositories with issue tracking enabled. They remain unavailable when none of the selected repositories has issue tracking enabled diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/MatchedPeopleToggle.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/MatchedPeopleToggle.tsx new file mode 100644 index 00000000000..5c7eaf22384 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/MatchedPeopleToggle.tsx @@ -0,0 +1,12 @@ +import { useId } from "react"; +import { Switch } from "@/components/ui/switch"; + +export function MatchedPeopleToggle({ checked, onChange }: { checked: boolean; onChange: (checked: boolean) => void }) { + const id = useId(); + return ( + + ); +} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx index 02d60a05029..84a4762f046 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ObservedDetails.tsx @@ -24,11 +24,16 @@ import { export function PullList({ pulls, provider, + matchedOnly = false, }: { + matchedOnly?: boolean; pulls: ObservedPull[]; provider: ObservedSnapshot["source_provider"]; }) { const terms = changeTerms(provider); + const emptyMessage = matchedOnly + ? "No merged changes from matched people in this period" + : `No ${terms.lower} in this period`; const [query, setQuery] = useState(""); const [limit, setLimit] = useState(20); const filtered = pulls.filter((pull) => @@ -104,7 +109,7 @@ export function PullList({

{filtered.length === 0 && (

- {query ? `No ${terms.lower} match this search` : `No ${terms.lower} in this period`} + {query ? `No ${terms.lower} match this search` : emptyMessage}

)}
@@ -150,9 +155,7 @@ export function PersonDetails({ {person.name} - - {person.email} · {person.logins.join(", ")} - + {[person.email, person.logins.join(", ")].filter(Boolean).join(" · ")} {onEdit && ( @@ -269,7 +281,7 @@ function PeopleTable({ {people.length === 0 && (
- {query ? `No engineers match “${query}”` : "Link accounts to see your engineers"} + {query ? `No engineers match “${query}”` : emptyMessage}
)}
@@ -380,8 +392,11 @@ export default function ObservedReport({ snapshot.people.length && snapshot.periods.current.merged_prs > 0 ? "people" : "pulls", ); const [accountEmail, setAccountEmail] = useState(null); - const [personEmail, setPersonEmail] = useState(null); - const person = snapshot.people.find((entry) => entry.email === personEmail) ?? null; + const [matchedOnly, setMatchedOnly] = useState(true); + const people = useMemo(() => reportPeople(snapshot, matchedOnly), [snapshot, matchedOnly]); + const pulls = useMemo(() => filterObservedPulls(snapshot, "current", matchedOnly), [snapshot, matchedOnly]); + const [personId, setPersonId] = useState(null); + const person = people.find((entry) => entry.id === personId) ?? null; const terms = changeTerms(snapshot.source_provider); const current = snapshot.periods.current; const baseline = snapshot.periods[comparison]; @@ -497,14 +512,15 @@ export default function ObservedReport({
- + - Engineers {snapshot.people.length} + Engineers {people.length} {terms.requests} Quality Branch spend + {activeTab !== "quality" && } {!readOnly && ( ' + "

This connection lasts 24 hours. Send disconnect in Slack to remove the saved session

", + ) + page.set_cookie(_cookie_name(token), csrf, max_age=600, secure=True, httponly=True, samesite="strict", path="/") + return page + + +@router.post(_PREFIX + "{token}", include_in_schema=False, response_class=HTMLResponse) +async def connect_account( + request: Request, + token: str, + context: Annotated[NativeAdminContext, Depends(native_admin_context)], +) -> Response: + base_url: Final = get_request_base_url(request) + parsed_base: Final = urlsplit(base_url) + origin: Final = f"{parsed_base.scheme}://{parsed_base.netloc}" + if parsed_base.scheme != "https" or request.headers.get("Origin") != origin: + raise HTTPException(403, "Reopen your private Slack connection link") + if request.headers.get("Content-Type", "").split(";", 1)[0] != "application/x-www-form-urlencoded": + raise HTTPException(400, "Expected a connection form") + form: Final = await request.form(max_fields=1, max_files=0, max_part_size=1024) + supplied: Final = form.get("csrf") + expected: Final = request.cookies.get(_cookie_name(token), "") + if ( + not isinstance(supplied, str) + or len(expected) != 43 + or len(supplied) != 43 + or not hmac.compare_digest(supplied.encode(), expected.encode()) + ): + raise HTTPException(403, "Reopen your private Slack connection link") + user_id: Final = await context.session_user(request) + if user_id is None: + raise HTTPException(401, "Your login expired. Reopen your private Slack connection link") + details: Final = await context.details(token) + user: Final = await context.admin(user_id, details) + await context.worker_request(token, context.mint_session(user)) + page: Final = _page( + "Account connected", "

Return to Slack and ask LiteAdmin to list your teams or check a budget

" + ) + page.delete_cookie(_cookie_name(token), path="/", secure=True, httponly=True, samesite="strict") + return page diff --git a/helm/litellm-helm/templates/deployment.yaml b/helm/litellm-helm/templates/deployment.yaml index cf7b3f8a38d..299d41e2019 100644 --- a/helm/litellm-helm/templates/deployment.yaml +++ b/helm/litellm-helm/templates/deployment.yaml @@ -57,6 +57,19 @@ spec: imagePullPolicy: {{ .Values.image.pullPolicy }} env: {{- include "litellm.proxyEnv" . | nindent 12 }} + {{- if .Values.liteadmin.enabled }} + - name: LITELLM_ADMIN_AGENT_URL + value: {{ printf "http://%s-liteadmin:10000" (include "litellm.fullname" . | trunc 53 | trimSuffix "-") | quote }} + - name: ADMIN_AGENT_SERVICE_TOKEN + valueFrom: + secretKeyRef: + name: {{ required "liteadmin.existingSecret is required" .Values.liteadmin.existingSecret }} + key: ADMIN_AGENT_SERVICE_TOKEN + {{- if not (hasKey (default dict .Values.envVars) "PROXY_BASE_URL") }} + - name: PROXY_BASE_URL + value: {{ required "liteadmin.gatewayUrl is required" .Values.liteadmin.gatewayUrl | quote }} + {{- end }} + {{- end }} {{- include "litellm.proxyMetricsEnv" . | nindent 12 }} {{- if .Values.collector.enabled }} {{- include "litellm.collectorEnv" . | nindent 12 }} diff --git a/helm/litellm-helm/templates/liteadmin.yaml b/helm/litellm-helm/templates/liteadmin.yaml new file mode 100644 index 00000000000..711edaf6913 --- /dev/null +++ b/helm/litellm-helm/templates/liteadmin.yaml @@ -0,0 +1,112 @@ +{{- if .Values.liteadmin.enabled }} +{{- $name := printf "%s-liteadmin" (include "litellm.fullname" . | trunc 53 | trimSuffix "-") }} +{{- $secret := required "liteadmin.existingSecret is required" .Values.liteadmin.existingSecret }} +apiVersion: apps/v1 +kind: Deployment +metadata: + name: {{ $name }} +spec: + replicas: 1 + strategy: + type: Recreate + selector: + matchLabels: + app.kubernetes.io/name: {{ $name }} + app.kubernetes.io/instance: {{ .Release.Name }} + template: + metadata: + labels: + app.kubernetes.io/name: {{ $name }} + app.kubernetes.io/instance: {{ .Release.Name }} + spec: + automountServiceAccountToken: false + terminationGracePeriodSeconds: 75 + {{- with .Values.imagePullSecrets }} + imagePullSecrets: + {{- toYaml . | nindent 8 }} + {{- end }} + securityContext: + runAsUser: 10001 + runAsGroup: 10001 + fsGroup: 10001 + runAsNonRoot: true + containers: + - name: liteadmin + image: "{{ .Values.image.repository }}:{{ .Values.image.tag | default .Chart.AppVersion }}" + imagePullPolicy: {{ .Values.image.pullPolicy }} + args: ["--admin-agent"] + securityContext: + allowPrivilegeEscalation: false + readOnlyRootFilesystem: true + capabilities: + drop: [ALL] + envFrom: + - secretRef: + name: {{ $secret }} + env: + - name: CONNECTION_AUTH_MODE + value: native + - name: LITELLM_BASE_URL + value: {{ required "liteadmin.gatewayUrl is required" .Values.liteadmin.gatewayUrl | quote }} + - name: LITELLM_MODEL + value: {{ required "liteadmin.model is required" .Values.liteadmin.model | quote }} + - name: STATE_DB + value: /var/data/events.sqlite3 + - name: ADMIN_READ_ONLY + value: {{ .Values.liteadmin.readOnly | quote }} + - name: OPENAI_AGENTS_DISABLE_TRACING + value: "1" + ports: + - name: health + containerPort: 10000 + readinessProbe: + httpGet: + path: /readyz + port: health + periodSeconds: 15 + livenessProbe: + httpGet: + path: /healthz + port: health + periodSeconds: 30 + resources: + {{- toYaml .Values.liteadmin.resources | nindent 12 }} + volumeMounts: + - name: state + mountPath: /var/data + - name: tmp + mountPath: /tmp + volumes: + - name: state + persistentVolumeClaim: + claimName: {{ $name }} + - name: tmp + emptyDir: + sizeLimit: 64Mi +--- +apiVersion: v1 +kind: Service +metadata: + name: {{ $name }} +spec: + type: ClusterIP + selector: + app.kubernetes.io/name: {{ $name }} + app.kubernetes.io/instance: {{ .Release.Name }} + ports: + - port: 10000 + targetPort: health +--- +apiVersion: v1 +kind: PersistentVolumeClaim +metadata: + name: {{ $name }} +spec: + accessModes: [ReadWriteOnce] + {{- with .Values.liteadmin.storageClassName }} + storageClassName: {{ . | quote }} + {{- end }} + resources: + requests: + storage: {{ .Values.liteadmin.storageSize }} +{{- end }} diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index 03d2a66a2b5..83dbb3c5aa0 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -3,6 +3,20 @@ # Declare variables to be passed into your templates. replicaCount: 1 +liteadmin: + enabled: false + existingSecret: "" + gatewayUrl: "" + model: "" + readOnly: false + storageSize: 1Gi + storageClassName: "" + resources: + requests: + cpu: 100m + memory: 256Mi + limits: + memory: 1Gi # numWorkers: 2 image: diff --git a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py index b51da626a60..5db9a92e51d 100644 --- a/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/byok_oauth_endpoints.py @@ -171,6 +171,42 @@ async def _session_key_is_live(session_key: str | None) -> bool: return True +async def get_authenticated_browser_user_id(request: Request) -> str | None: + from datetime import datetime, timezone + + from pydantic import TypeAdapter, ValidationError + + from litellm.proxy._types import hash_token + from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_key_object + from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache + + user_id, session_key = _session_identity_from_cookie(request) + if not user_id or not session_key or prisma_client is None: + return None + try: + auth: Final = ( + await get_key_object( + hash_token(session_key), + prisma_client, + user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + check_db_only=True, + ) + if session_key.startswith("sk-") + else ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(session_key) + ) + except Exception: + return None + if auth is None or auth.user_id != user_id or auth.blocked or auth.expires is None: + return None + try: + expiration: Final = TypeAdapter(datetime).validate_python(auth.expires) + except ValidationError: + return None + expires: Final = expiration.replace(tzinfo=timezone.utc) if expiration.tzinfo is None else expiration + return user_id if expires > datetime.now(timezone.utc) else None + + async def _byok_session_auth(request: Request) -> UserAPIKeyAuth: """Require the UI session cookie, with the embedded session key re-resolved against the DB so a revoked (logged-out) session cannot diff --git a/tests/unit/enterprise/proxy/test_liteadmin.py b/tests/unit/enterprise/proxy/test_liteadmin.py new file mode 100644 index 00000000000..4d1c3ce7f0a --- /dev/null +++ b/tests/unit/enterprise/proxy/test_liteadmin.py @@ -0,0 +1,244 @@ +from __future__ import annotations + +import json +import re +from typing import Final + +import httpx +import pytest +from fastapi import FastAPI, Request +from fastapi.testclient import TestClient +from litellm_enterprise.proxy.liteadmin import AdminSession, NativeAdminContext, native_admin_context, router +from pydantic import SecretStr, TypeAdapter + +from litellm.proxy._types import LiteLLM_UserTable + +TOKEN: Final = "a" * 43 +PATH: Final = "/liteadmin/slack/connect/" + TOKEN +ORIGIN: Final = "https://gateway.example.com" + + +class Worker: + def __init__(self, *, email: str = "alice@example.com", status: int = 200) -> None: + self.email = email + self.status = status + self.session: object = None + self.role = "proxy_admin" + + def request(self, request: httpx.Request) -> httpx.Response: + assert request.headers["X-LiteLLM-Admin-Agent-Token"] == "s" * 32 + if request.method == "POST": + self.session = json.loads(request.content) + return httpx.Response(self.status, json={"status": "connected"}) + return httpx.Response( + self.status, + json={ + "workspace_id": "Tworkspace", + "slack_user_id": "Ualice", + "email": self.email, + }, + ) + + +def client_for( + worker: Worker, *, role: str | None = None, logged_in: bool = True, email: str = "alice@example.com" +) -> TestClient: + async def session_user(request: Request) -> str | None: + return "alice" if logged_in else None + + async def load_user(user_id: str) -> LiteLLM_UserTable: + return LiteLLM_UserTable(user_id=user_id, user_email=email, user_role=role or worker.role) + + def mint(user: LiteLLM_UserTable) -> AdminSession: + return AdminSession(user_id=user.user_id, credential=SecretStr("personal-session"), expires_at=86400) + + context: Final = NativeAdminContext( + "http://private-worker:10000", + SecretStr("s" * 32), + httpx.AsyncClient(transport=httpx.MockTransport(worker.request)), + session_user, + load_user, + mint, + ) + app: Final = FastAPI() + app.include_router(router) + app.dependency_overrides[native_admin_context] = lambda: context + return TestClient(app, base_url=ORIGIN) + + +def csrf_from(client: TestClient) -> str: + page: Final = client.get(PATH) + assert page.status_code == 200 + match: Final = re.search('name="csrf" value="([^"]+)"', page.text) + assert match is not None + return match[1] + + +def test_connect_uses_existing_login_without_a_hosted_oauth_callback(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PROXY_BASE_URL", ORIGIN) + with client_for(Worker(), logged_in=False) as client: + response: Final = client.get(PATH, follow_redirects=False) + assert response.status_code == 303 + assert ( + response.headers["location"] == ORIGIN + "/sso/key/generate?return_to=%2Fliteadmin%2Fslack%2Fconnect%2F" + TOKEN + ) + assert response.headers["cache-control"] == "no-store" + + +def test_connect_hands_off_personal_session_only_over_private_worker_channel(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PROXY_BASE_URL", ORIGIN) + worker: Final = Worker() + with client_for(worker) as client: + csrf: Final = csrf_from(client) + response: Final = client.post(PATH, data={"csrf": csrf}, headers={"Origin": ORIGIN}) + assert response.status_code == 200 + assert "Account connected" in response.text + assert worker.session == {"user_id": "alice", "credential": "personal-session", "expires_at": 86400.0} + assert "personal-session" not in response.text + assert "Max-Age=0" in response.headers["set-cookie"] + + +@pytest.mark.parametrize("role,email", [("internal_user", "alice@example.com"), ("proxy_admin", "bob@example.com")]) +def test_connect_rejects_nonadmin_and_another_slack_users_link( + monkeypatch: pytest.MonkeyPatch, + role: str, + email: str, +) -> None: + monkeypatch.setenv("PROXY_BASE_URL", ORIGIN) + worker: Final = Worker(email=email) + with client_for(worker, role=role) as client: + response: Final = client.get(PATH) + assert response.status_code == 403 + assert worker.session is None + + +@pytest.mark.parametrize("origin,csrf", [("https://attacker.example", None), ("null", None), (ORIGIN, "b" * 43)]) +def test_connect_requires_same_origin_and_browser_csrf( + monkeypatch: pytest.MonkeyPatch, + origin: str, + csrf: str | None, +) -> None: + monkeypatch.setenv("PROXY_BASE_URL", ORIGIN) + worker: Final = Worker() + with client_for(worker) as client: + valid: Final = csrf_from(client) + response: Final = client.post(PATH, data={"csrf": csrf or valid}, headers={"Origin": origin}) + assert response.status_code == 403 + assert worker.session is None + + +@pytest.mark.parametrize("status,expected", [(410, 410), (403, 403), (500, 503), (302, 503)]) +def test_worker_denial_expiry_and_failure_never_mint_a_session( + monkeypatch: pytest.MonkeyPatch, + status: int, + expected: int, +) -> None: + monkeypatch.setenv("PROXY_BASE_URL", ORIGIN) + worker: Final = Worker(status=status) + with client_for(worker) as client: + response: Final = client.get(PATH) + assert response.status_code == expected + assert worker.session is None + + +def test_csrf_cookie_cannot_cross_links(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PROXY_BASE_URL", ORIGIN) + worker: Final = Worker() + with client_for(worker) as client: + csrf: Final = csrf_from(client) + response: Final = client.post(PATH.replace(TOKEN, "b" * 43), data={"csrf": csrf}, headers={"Origin": ORIGIN}) + assert response.status_code == 403 + assert worker.session is None + + +@pytest.mark.parametrize( + "url,secret,enterprise,database,status", + [ + ("", "s" * 32, True, True, 404), + ("http://worker:10000", "s" * 32, False, True, 403), + ("http://worker:10000", "s" * 32, True, False, 503), + ("file:///etc/passwd", "s" * 32, True, True, 503), + ("https://user:password@worker", "s" * 32, True, True, 503), + ("https://worker/path", "s" * 32, True, True, 503), + ("https://worker", "short", True, True, 503), + ("http://[broken", "s" * 32, True, True, 503), + ("http://worker:broken", "s" * 32, True, True, 503), + ], +) +def test_native_configuration_requires_enterprise_database_and_private_worker_credentials( + url: str, + secret: str, + enterprise: bool, + database: bool, + status: int, +) -> None: + from fastapi import HTTPException + from litellm_enterprise.proxy.liteadmin import validate_native_configuration + + with pytest.raises(HTTPException) as error: + validate_native_configuration(url, secret, enterprise, database) + assert error.value.status_code == status + + +def test_connect_rechecks_admin_permission_after_consent_page(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PROXY_BASE_URL", ORIGIN) + worker: Final = Worker() + with client_for(worker) as client: + csrf: Final = csrf_from(client) + worker.role = "internal_user" + response: Final = client.post(PATH, data={"csrf": csrf}, headers={"Origin": ORIGIN}) + assert response.status_code == 403 + assert worker.session is None + + +def test_consent_escapes_slack_email(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("PROXY_BASE_URL", ORIGIN) + email: Final = '@example.com' + with client_for(Worker(email=email), email=email) as client: + page: Final = client.get(PATH) + assert page.status_code == 200 + assert "