diff --git a/.github/ci-coverage-allowlist.yml b/.github/ci-coverage-allowlist.yml index 4432c19bac6..ff8fa864d4a 100644 --- a/.github/ci-coverage-allowlist.yml +++ b/.github/ci-coverage-allowlist.yml @@ -49,18 +49,15 @@ test_paths: paths: - tests/code_coverage_tests/test_aio_http_image_conversion.py - reason: >- - What is left of a second mirror that sat beside tests/test_litellm and ran nowhere. Its - other 30 files moved into the real mirror on 2026-08-20 and now run; these four cannot, - because each shares a filename with a live test whose contents are disjoint from it, so - landing them means merging test bodies rather than moving a file. Measured on the same - date: test_common_utils.py holds 15 tests the live file does not, test_oci_chat_transformation - 13, test_deepseek_chat_transformation 12, and test_discoverable_endpoints 5. Revisit by - merging each into its twin, which is a content review, not a move + The last file of a second mirror that sat beside tests/test_litellm and ran nowhere. Its + other 33 files landed in the real mirror during August 2026, 30 as moves and 3 by merging + their bodies into the live file of the same name. This one cannot follow either route yet: + its live twin was rewritten from 1268 lines to 9434, and of the 19 tests here 5 have no + counterpart while 25 assertions fail against today's code, so what survives that rewrite + is a judgement about the endpoints, not a merge. Revisit by deciding which of the five + behaviours still hold paths: - - tests/litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py - - tests/litellm/llms/oci/chat/test_oci_chat_transformation.py - tests/litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py - - tests/litellm/proxy/management_endpoints/test_common_utils.py - reason: >- No job invokes this suite and its files mix pure transformation tests with ones driving live vendor vector stores, so assigning them needs a per-file decision diff --git a/.github/scripts/assert_ci_coverage.py b/.github/scripts/assert_ci_coverage.py index 86a6b7d4e72..c8572d9f6ef 100644 --- a/.github/scripts/assert_ci_coverage.py +++ b/.github/scripts/assert_ci_coverage.py @@ -1,11 +1,12 @@ from __future__ import annotations import ast +import operator import pathlib import re import sys import warnings -from collections.abc import Iterable, Mapping, Sequence +from collections.abc import Callable, Iterable, Mapping, Sequence from dataclasses import dataclass from typing import Final @@ -56,6 +57,14 @@ class Allowlist: return any(relative_path == path for entry in self.dockerfiles for path in entry.paths) +@dataclass(frozen=True, slots=True) +class Section: + name: str + entries: tuple[AllowEntry, ...] + candidates: tuple[str, ...] + matches: Callable[[str, str], bool] + + @dataclass(frozen=True, slots=True) class Scalar: key: str @@ -368,6 +377,25 @@ def _uncovered_dockerfiles(allowlist: Allowlist, tokens: frozenset[str]) -> tupl ) +def _stale_allowlist_paths( + allowlist: Allowlist, + *, + test_files: tuple[str, ...], + dockerfiles: tuple[str, ...], +) -> tuple[Finding, ...]: + sections: Final[tuple[Section, ...]] = ( + Section("test_paths", allowlist.test_paths, test_files, _token_covers), + Section("dockerfiles", allowlist.dockerfiles, dockerfiles, operator.eq), + ) + return tuple( + Finding(subject=path, detail=f"listed under '{section.name}' but matches no file the census looks at") + for section in sections + for entry in section.entries + for path in entry.paths + if not any(section.matches(path, candidate) for candidate in section.candidates) + ) + + def _parse_entry(item: object, section: str) -> AllowEntry: if not isinstance(item, dict): raise SystemExit(f"{ALLOWLIST_FILE.name}: '{section}' entries must be mappings") @@ -465,7 +493,14 @@ def main() -> int: test_findings = _uncovered_tests(allowlist, _invoked_test_tokens(scalars)) dockerfile_findings = _uncovered_dockerfiles(allowlist, _built_dockerfile_tokens(scalars)) + stale_findings = _stale_allowlist_paths(allowlist, test_files=_test_files(), dockerfiles=_dockerfiles()) + if stale_findings: + _report( + "allowlist entries that exempt nothing", + stale_findings, + "Delete each from .github/ci-coverage-allowlist.yml; the file it named is gone or was renamed.", + ) if test_findings: _report( "test files that no CI job invokes", @@ -478,7 +513,7 @@ def main() -> int: dockerfile_findings, "Build each in a workflow, or list it in .github/ci-coverage-allowlist.yml with a reason.", ) - if test_findings or dockerfile_findings: + if stale_findings or test_findings or dockerfile_findings: return 1 _write( diff --git a/.github/scripts/assert_workflow_dir_hygiene.py b/.github/scripts/assert_workflow_dir_hygiene.py new file mode 100644 index 00000000000..681a365b1ba --- /dev/null +++ b/.github/scripts/assert_workflow_dir_hygiene.py @@ -0,0 +1,149 @@ +#!/usr/bin/env python3 +"""Three invariants about what lives in .github/workflows/ and what its names mean. + +`.github/workflows/` is a directory GitHub reads, not a place to keep things. Every +file at its top level is parsed as a workflow, so a script or a data file parked there +is either an invalid workflow or an orphan nobody can find. A subdirectory is not read +at all, so helper files may live in one. GitHub accepts both `.yml` and `.yaml`, and +this repo spells them `.yml`, which is a naming rule rather than a validity one and is +reported separately. And the `_` prefix is the repo's only signal that a workflow is a +reusable building block rather than something that runs on its own, which is worth +nothing unless it is true both ways. + + WF001 a top-level file in .github/workflows/ that is not a workflow at all + WF002 a workflow whose only trigger is `workflow_call` but is not `_`-prefixed + WF003 a `_`-prefixed workflow that no other workflow can call + WF004 a real workflow spelled `.yaml` where this directory spells them `.yml` + +A workflow with `workflow_call` alongside a human trigger is deliberately dual-mode +and belongs under its plain name, so only the call-only ones are held to WF002. + +Usage +----- + python assert_workflow_dir_hygiene.py + +Exit code 1 if any violation is found. +""" + +from __future__ import annotations + +import pathlib +import sys +from dataclasses import dataclass +from typing import Final + +import yaml + +REPO_ROOT: Final = pathlib.Path(__file__).resolve().parents[2] +WORKFLOW_DIR: Final = REPO_ROOT / ".github" / "workflows" +SCRIPT_HOME: Final = ".github/scripts/" +REUSABLE_PREFIX: Final = "_" +CALL_TRIGGER: Final = "workflow_call" +CANONICAL_SUFFIX: Final = ".yml" +WORKFLOW_SUFFIXES: Final = frozenset((CANONICAL_SUFFIX, ".yaml")) + + +@dataclass(frozen=True, slots=True) +class Finding: + subject: str + code: str + detail: str + + def render(self) -> str: + return f" - {self.subject}: {self.code} {self.detail}" + + +def _triggers(document: object) -> frozenset[str]: + if not isinstance(document, dict): + return frozenset() + raw: Final = document.get("on", document.get(True)) + if isinstance(raw, str): + return frozenset({raw}) + if isinstance(raw, dict): + return frozenset(str(key) for key in raw) + if isinstance(raw, list): + return frozenset(str(item) for item in raw) + return frozenset() + + +def _workflows(directory: pathlib.Path) -> tuple[pathlib.Path, ...]: + return tuple( + path + for path in sorted(directory.iterdir()) + if path.is_file() and path.suffix in WORKFLOW_SUFFIXES + ) + + +def _strays(directory: pathlib.Path) -> tuple[Finding, ...]: + return tuple( + Finding( + path.name, + "WF001", + f"is not a workflow, and GitHub parses every top-level file here as one; " + f"move it to {SCRIPT_HOME} or into a subdirectory, which GitHub does not read", + ) + for path in sorted(directory.iterdir()) + if path.is_file() and path.suffix not in WORKFLOW_SUFFIXES + ) + + +def _misspelled(directory: pathlib.Path) -> tuple[Finding, ...]: + return tuple( + Finding( + path.name, + "WF004", + f"is a real workflow and GitHub reads it, but this directory spells them " + f"{CANONICAL_SUFFIX}; rename it to {path.stem}{CANONICAL_SUFFIX}", + ) + for path in _workflows(directory) + if path.suffix != CANONICAL_SUFFIX + ) + + +def _misnamed(directory: pathlib.Path) -> tuple[Finding, ...]: + return tuple( + finding + for path in _workflows(directory) + for finding in _naming_findings(path, _triggers(yaml.safe_load(path.read_text(encoding="utf-8")))) + ) + + +def _naming_findings(path: pathlib.Path, triggers: frozenset[str]) -> tuple[Finding, ...]: + underscored: Final = path.name.startswith(REUSABLE_PREFIX) + if triggers == frozenset({CALL_TRIGGER}) and not underscored: + return ( + Finding( + path.name, + "WF002", + f"is only callable by another workflow, so name it {REUSABLE_PREFIX}{path.name}", + ), + ) + if underscored and CALL_TRIGGER not in triggers: + return ( + Finding( + path.name, + "WF003", + f"is named as a reusable workflow but has no {CALL_TRIGGER} trigger; " + "add one or drop the prefix", + ), + ) + return () + + +def main() -> int: + findings: Final = _strays(WORKFLOW_DIR) + _misspelled(WORKFLOW_DIR) + _misnamed(WORKFLOW_DIR) + if not findings: + total: Final = len(_workflows(WORKFLOW_DIR)) + sys.stdout.write( + f"OK: {total} workflows, every file in .github/workflows/ is one, and the " + f"{REUSABLE_PREFIX} prefix means callable in both directions.\n" + ) + return 0 + sys.stdout.write("ERROR: .github/workflows/ holds files that break its own conventions\n") + for finding in findings: + sys.stdout.write(f"{finding.render()}\n") + return 1 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/.github/workflows/_test-unit-base.yml b/.github/workflows/_test-unit-base.yml index 4f4339a360a..54f50524a39 100644 --- a/.github/workflows/_test-unit-base.yml +++ b/.github/workflows/_test-unit-base.yml @@ -129,6 +129,13 @@ jobs: WORKERS: ${{ inputs.workers }} RERUNS: ${{ inputs.reruns }} DIST: ${{ inputs.dist }} + # coverage.py's sys.monitoring backend (PEP 669), the cheapest core it has. + # It is only the default from Python 3.14, and these shards run 3.12, so it + # has to be asked for. Coverage refuses it when branch measurement is on + # (`branch_right_left` needs > 3.14.0a5) and falls back to the slow core with + # a `no-sysmon` warning, so turning on `branch = true` here means giving this + # back until the runners move to 3.14. + COVERAGE_CORE: sysmon run: | if [ "${WORKERS}" = "0" ]; then uv run --no-sync pytest ${TEST_PATH:?} \ diff --git a/.github/workflows/ci-coverage.yml b/.github/workflows/ci-coverage.yml index 486587fc27d..7bc476db134 100644 --- a/.github/workflows/ci-coverage.yml +++ b/.github/workflows/ci-coverage.yml @@ -46,3 +46,6 @@ jobs: # nowhere while counting as covered, which is how the caching suite went unrun. - name: Assert no -k expression deselects a file from every job that globs it run: python .github/scripts/assert_ci_coverage.py --slices + + - name: Assert .github/workflows/ holds only workflows, correctly named + run: python .github/scripts/assert_workflow_dir_hygiene.py diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index 5a180c13c53..e031ba46773 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -122,6 +122,11 @@ jobs: uv run --no-sync ruff check . cd .. + - name: Run Ruff linting (test tree) + if: steps.changes.outputs.decision != 'skip' + run: | + uv run --no-sync ruff check --config ruff-tests.toml tests + - name: Check strict-rule budget (delta vs base) if: steps.changes.outputs.decision != 'skip' run: | @@ -132,7 +137,7 @@ jobs: run: | uv run --no-sync python scripts/type_discipline_gate.py --base "$GATE_BASE_SHA" - - name: Check test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes, litellm global mutation, credential-gated skips, delta vs base) + - name: Check test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes, litellm global mutation, credential-gated skips, conftest snapshot inventory, delta vs base) if: steps.changes.outputs.decision != 'skip' run: | uv run --no-sync python scripts/test_quality_gate.py --base "$GATE_BASE_SHA" diff --git a/.github/workflows/test-model-map.yaml b/.github/workflows/test-model-map.yml similarity index 100% rename from .github/workflows/test-model-map.yaml rename to .github/workflows/test-model-map.yml diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index fbba9969c28..3d6fffe7304 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -113,6 +113,7 @@ jobs: tests/test_litellm/rag tests/test_litellm/realtime_api tests/test_litellm/rerank_api + tests/test_litellm/rust_bridge tests/test_litellm/sandbox tests/test_litellm/test_router tests/test_litellm/vector_stores diff --git a/Makefile b/Makefile index 5ae2638fbaa..e17fdba3c85 100644 --- a/Makefile +++ b/Makefile @@ -160,6 +160,7 @@ lint-format-check-changed: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) # Linting targets lint-ruff: $(LINT_DEP_INSTALL) cd litellm && $(UV_RUN) ruff check . && cd .. + $(UV_RUN) ruff check --config ruff-tests.toml tests # faster linter for developing ... # inspiration from: @@ -205,8 +206,8 @@ lint-type-discipline: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) $(UV_RUN) python scripts/type_discipline_gate.py --base origin/litellm_internal_staging # Test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes, -# litellm module-global mutation, credential-gated skips), counted across tests/ the -# same delta-vs-base way. +# litellm module-global mutation, credential-gated skips, conftest snapshot +# inventory), counted across tests/ the same delta-vs-base way. lint-test-quality: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE) $(UV_RUN) python scripts/test_quality_gate.py --base origin/litellm_internal_staging diff --git a/basedpyright-code-budget.json b/basedpyright-code-budget.json index b4c324a2c4c..776aecbd883 100644 --- a/basedpyright-code-budget.json +++ b/basedpyright-code-budget.json @@ -105,13 +105,13 @@ "limit": 109 }, "reportUnknownMemberType": { - "limit": 39017 + "limit": 39011 }, "reportUnknownParameterType": { "limit": 19885 }, "reportUnknownVariableType": { - "limit": 30572 + "limit": 30569 }, "reportUnnecessaryCast": { "limit": 117 diff --git a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py index a8e46349917..4bb00408fc3 100644 --- a/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py +++ b/enterprise/litellm_enterprise/proxy/common_utils/check_batch_cost.py @@ -255,6 +255,52 @@ class CheckBatchCost: "so it will no longer be polled" ) + async def _claim_job_for_costing(self, job: "LiteLLM_ManagedObjectTable") -> bool: + """ + Atomically flip batch_processed from false to true, returning whether this pod won + the row. Every pod and uvicorn worker schedules its own poller against the shared + table, so without this compare-and-swap two of them can select the same completed + batch in one window and both emit an aretrieve_batch spend log for it. Schemas + without the column can't be claimed, so they keep the pre-existing behavior. + + Called immediately before the spend log is written rather than before the results + fetch, because batch_processed is also what holds off deletion of the files that + fetch reads and what keeps an unbilled row selectable by the next poll cycle. + """ + if not self._has_batch_processed_column: + return True + try: + claimed: Final = await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={"id": job.id, "batch_processed": False}, + data={"batch_processed": True}, + ) + except Exception as db_err: + verbose_proxy_logger.error( + f"CheckBatchCost: failed to claim job {job.id} for cost tracking: {db_err}" + ) + return False + return claimed > 0 + + async def _release_job_claim(self, job: "LiteLLM_ManagedObjectTable") -> None: + """Give a claimed row back once billing it failed, so a later poll cycle retries it. + + Safe to match on batch_processed=True: while this poller is active the retrieve + path leaves the column alone (batch_cost_poller_is_active), so a true value here + is always this pod's own claim. + """ + if not self._has_batch_processed_column: + return + try: + await self.prisma_client.db.litellm_managedobjecttable.update_many( + where={"id": job.id, "batch_processed": True}, + data={"batch_processed": False}, + ) + except Exception as db_err: + verbose_proxy_logger.error( + f"CheckBatchCost: failed to release the claim on job {job.id}, " + f"so its cost will not be retried: {db_err}" + ) + @staticmethod def _has_unified_id_without_model(job: "LiteLLM_ManagedObjectTable") -> bool: """A unified id that decodes but carries no model_id can never be routed.""" @@ -572,9 +618,10 @@ class CheckBatchCost: """ Fetch a completed batch's results, compute cost/usage, and emit the aretrieve_batch spend log. Returns (model_name, llm_provider) on - success, None when the job can't be routed to a deployment. Raises on - results-fetch or cost-computation failures so the caller can leave the - job unprocessed and retry it on a later poll. + success, None when the job can't be routed to a deployment or when + another pod claimed it. Raises on results-fetch or cost-computation + failures so the caller can leave the job unprocessed and retry it on a + later poll. """ from litellm.batches.batch_utils import ( _get_file_content_as_dictionary, @@ -743,12 +790,23 @@ class CheckBatchCost: optional_params={}, ) - await logging_obj.async_success_handler( - result=response, - batch_cost=batch_cost, - batch_usage=batch_usage, - batch_models=batch_models, - ) + if not await self._claim_job_for_costing(job): + verbose_proxy_logger.info( + f"CheckBatchCost: batch {batch_id} (job {job.id}) was claimed by another pod " + "in this window, so its cost is already being tracked there" + ) + return None + + try: + await logging_obj.async_success_handler( + result=response, + batch_cost=batch_cost, + batch_usage=batch_usage, + batch_models=batch_models, + ) + except Exception: + await self._release_job_claim(job) + raise # Record batch duration (completed_at - created_at) if prom_logger and response.completed_at and response.created_at: diff --git a/enterprise/pyproject.toml b/enterprise/pyproject.toml index bb580c82760..8bbde7f3764 100644 --- a/enterprise/pyproject.toml +++ b/enterprise/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-enterprise" -version = "0.1.57" +version = "0.1.58" description = "Package for LiteLLM Enterprise features" readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.1.57" +version = "0.1.58" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-enterprise==", diff --git a/helm/litellm-helm/templates/deployment.yaml b/helm/litellm-helm/templates/deployment.yaml index 32bfa4b2647..52ffd117535 100644 --- a/helm/litellm-helm/templates/deployment.yaml +++ b/helm/litellm-helm/templates/deployment.yaml @@ -100,6 +100,13 @@ spec: - name: DATABASE_URL value: {{ .Values.db.url | quote }} {{- end }} + {{- if and .Values.db.useExisting .Values.db.readReplicaUrl .Values.db.secret.readReplicaEndpointKey (not .Values.db.secret.readReplicaUrlKey) }} + - name: DATABASE_READER_HOST + valueFrom: + secretKeyRef: + name: {{ .Values.db.secret.name }} + key: {{ .Values.db.secret.readReplicaEndpointKey }} + {{- end }} {{- if and .Values.db.useExisting .Values.db.secret.readReplicaUrlKey }} - name: DATABASE_URL_READ_REPLICA valueFrom: diff --git a/helm/litellm-helm/tests/deployment_tests.yaml b/helm/litellm-helm/tests/deployment_tests.yaml index b11c445889e..ee946038202 100644 --- a/helm/litellm-helm/tests/deployment_tests.yaml +++ b/helm/litellm-helm/tests/deployment_tests.yaml @@ -80,6 +80,96 @@ tests: secretKeyRef: name: my-secret key: my-key + - it: should inject DATABASE_READER_HOST from readReplicaEndpointKey before DATABASE_URL_READ_REPLICA + template: deployment.yaml + set: + db: + deployStandalone: false + useExisting: true + secret: + name: postgres + usernameKey: username + passwordKey: password + readReplicaEndpointKey: reader-host + readReplicaUrl: postgresql://$(DATABASE_USERNAME):$(DATABASE_PASSWORD)@$(DATABASE_READER_HOST):5432/$(DATABASE_NAME)?sslmode=require + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_READER_HOST + valueFrom: + secretKeyRef: + name: postgres + key: reader-host + - contains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_URL_READ_REPLICA + value: postgresql://$(DATABASE_USERNAME):$(DATABASE_PASSWORD)@$(DATABASE_READER_HOST):5432/$(DATABASE_NAME)?sslmode=require + # $(VAR) interpolation only resolves vars defined EARLIER in the env + # array, so the reader host must precede the composed URL + - equal: + path: spec.template.spec.containers[0].env[7].name + value: DATABASE_READER_HOST + - equal: + path: spec.template.spec.containers[0].env[8].name + value: DATABASE_URL_READ_REPLICA + - it: should omit reader host when readReplicaUrl is unset + template: deployment.yaml + set: + db: + deployStandalone: false + useExisting: true + secret: + name: postgres + usernameKey: username + passwordKey: password + readReplicaEndpointKey: reader-host + asserts: + - notContains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_READER_HOST + valueFrom: + secretKeyRef: + name: postgres + key: reader-host + - it: should prefer readReplicaUrlKey over readReplicaEndpointKey composition + template: deployment.yaml + set: + db: + useExisting: true + secret: + name: postgres + usernameKey: username + passwordKey: password + readReplicaUrlKey: reader-url + readReplicaEndpointKey: reader-host + readReplicaUrl: postgresql://ignored + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_URL_READ_REPLICA + valueFrom: + secretKeyRef: + name: postgres + key: reader-url + - notContains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_URL_READ_REPLICA + value: postgresql://ignored + # the unused reader-host secret ref must be suppressed so a missing + # key can't fail pod creation + - notContains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_READER_HOST + valueFrom: + secretKeyRef: + name: postgres + key: reader-host - it: should work with extraEnvVars template: deployment.yaml set: diff --git a/helm/litellm-helm/values.yaml b/helm/litellm-helm/values.yaml index 4ef8fc97b27..f8df98de102 100644 --- a/helm/litellm-helm/values.yaml +++ b/helm/litellm-helm/values.yaml @@ -277,6 +277,14 @@ db: # written to db.readReplicaUrl ends up visible in the rendered pod spec # and the Helm release secret. readReplicaUrlKey: "" + # Optional: when set, a DATABASE_READER_HOST env var is sourced from this + # secret key, so db.readReplicaUrl can compose the reader URL from + # individual secret components, e.g. + # postgresql://$(DATABASE_USERNAME):$(DATABASE_PASSWORD)@$(DATABASE_READER_HOST):5432/$(DATABASE_NAME) + # Use this when your secret store holds the bare reader hostname rather + # than a full connection URL. Only takes effect when readReplicaUrl is + # set; ignored when readReplicaUrlKey is set. + readReplicaEndpointKey: "" # Optional read-replica routing. When set, the proxy sends read-only # queries (find_*, count, group_by, query_raw/_first) to this URL while diff --git a/helm/litellm/templates/_helpers.tpl b/helm/litellm/templates/_helpers.tpl index bffd627393a..72f7f74bcf6 100644 --- a/helm/litellm/templates/_helpers.tpl +++ b/helm/litellm/templates/_helpers.tpl @@ -213,18 +213,21 @@ whenever the password contains a URL-reserved character (@, /, ?, %, +, When `database.writer.useIAMAuth: true`, the chart injects IAM_TOKEN_DB_AUTH=true and omits DATABASE_PASSWORD — the entrypoint mints -the URL from DATABASE_HOST/PORT/USER/NAME plus a short-lived IAM token -instead of a static password. +the URL from DATABASE_HOST/PORT/USER/NAME plus a short-lived AWS RDS IAM +token instead of a static password. `database.writer.useAzureEntraAuth: true` +does the same with AZURE_POSTGRESQL_AUTH=true and a Microsoft Entra ID token, +for Azure Database for PostgreSQL. The two are mutually exclusive. The read replica is opt-in via `database.reader.host`. The chart emits DATABASE_HOST_READ_REPLICA / DATABASE_PORT_READ_REPLICA / DATABASE_NAME_READ_REPLICA (+ DATABASE_SCHEMA_READ_REPLICA) for both auth modes, plus DATABASE_USER_READ_REPLICA / DATABASE_PASSWORD_READ_REPLICA for -password auth. When `database.reader.useIAMAuth: true` it omits +password auth. When `database.reader.useIAMAuth: true` (or +`database.reader.useAzureEntraAuth: true`) it omits DATABASE_PASSWORD_READ_REPLICA and the entrypoint mints the reader URL the -same way. Reader IAM only takes effect when the writer also uses IAM auth -(the proxy gates URL minting on IAM_TOKEN_DB_AUTH, which only the writer -sets). +same way. Reader token auth only takes effect when the writer uses the same +token source, since the proxy gates URL minting on the single global +IAM_TOKEN_DB_AUTH / AZURE_POSTGRESQL_AUTH toggle that only the writer sets. */}} {{- define "litellm.serverEnv" -}} {{- $root := .root -}} @@ -254,9 +257,15 @@ sets). - name: DATABASE_SCHEMA value: {{ .schema | quote }} {{- end }} +{{- if and .useIAMAuth .useAzureEntraAuth }} +{{- fail "database.writer.useIAMAuth and database.writer.useAzureEntraAuth are mutually exclusive: the database password can only come from one token source" }} +{{- end }} {{- if .useIAMAuth }} - name: IAM_TOKEN_DB_AUTH value: "true" +{{- else if .useAzureEntraAuth }} +- name: AZURE_POSTGRESQL_AUTH + value: "true" {{- else }} - name: DATABASE_PASSWORD valueFrom: @@ -270,6 +279,9 @@ sets). {{- if and .useIAMAuth (not $root.Values.database.writer.useIAMAuth) }} {{- fail "database.reader.useIAMAuth requires database.writer.useIAMAuth: true (the proxy gates IAM URL minting on IAM_TOKEN_DB_AUTH, which is only set by the writer)" }} {{- end }} +{{- if and .useAzureEntraAuth (not $root.Values.database.writer.useAzureEntraAuth) }} +{{- fail "database.reader.useAzureEntraAuth requires database.writer.useAzureEntraAuth: true (the proxy gates Entra URL minting on AZURE_POSTGRESQL_AUTH, which is only set by the writer)" }} +{{- end }} - name: DATABASE_HOST_READ_REPLICA value: {{ .host | quote }} - name: DATABASE_PORT_READ_REPLICA @@ -280,7 +292,7 @@ sets). - name: DATABASE_SCHEMA_READ_REPLICA value: {{ .schema | quote }} {{- end }} -{{- if .useIAMAuth }} +{{- if or .useIAMAuth .useAzureEntraAuth }} {{- if .passwordSecret.name }} - name: DATABASE_USER_READ_REPLICA valueFrom: diff --git a/helm/litellm/tests/database_auth_tests.yaml b/helm/litellm/tests/database_auth_tests.yaml new file mode 100644 index 00000000000..adbe14c59c2 --- /dev/null +++ b/helm/litellm/tests/database_auth_tests.yaml @@ -0,0 +1,116 @@ +suite: test database token auth env vars +templates: + - gateway/deployment.yaml + - gateway/configmap.yaml + - backend/deployment.yaml + - backend/configmap.yaml +values: + - ./values/required.yaml +tests: + - it: writer emits DATABASE_PASSWORD and no token toggle by default + template: gateway/deployment.yaml + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_PASSWORD + valueFrom: + secretKeyRef: + name: litellm-writer-secret + key: password + any: true + - notContains: + path: spec.template.spec.containers[0].env + content: + name: IAM_TOKEN_DB_AUTH + value: "true" + any: true + - notContains: + path: spec.template.spec.containers[0].env + content: + name: AZURE_POSTGRESQL_AUTH + value: "true" + any: true + + - it: writer emits AZURE_POSTGRESQL_AUTH and omits DATABASE_PASSWORD under Entra auth + template: gateway/deployment.yaml + set: + database.writer.useAzureEntraAuth: true + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: AZURE_POSTGRESQL_AUTH + value: "true" + any: true + - notContains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_PASSWORD + any: true + - notContains: + path: spec.template.spec.containers[0].env + content: + name: IAM_TOKEN_DB_AUTH + value: "true" + any: true + + - it: backend gets the same Entra toggle as the gateway + template: backend/deployment.yaml + set: + database.writer.useAzureEntraAuth: true + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: AZURE_POSTGRESQL_AUTH + value: "true" + any: true + + - it: writer rejects both token sources at once + template: gateway/deployment.yaml + set: + database.writer.useIAMAuth: true + database.writer.useAzureEntraAuth: true + asserts: + - failedTemplate: + errorMessage: "database.writer.useIAMAuth and database.writer.useAzureEntraAuth are mutually exclusive: the database password can only come from one token source" + + - it: reader Entra auth without writer Entra auth is rejected + template: gateway/deployment.yaml + set: + database.reader.host: reader.example.com + database.reader.dbname: litellm + database.reader.useAzureEntraAuth: true + asserts: + - failedTemplate: + errorMessage: "database.reader.useAzureEntraAuth requires database.writer.useAzureEntraAuth: true (the proxy gates Entra URL minting on AZURE_POSTGRESQL_AUTH, which is only set by the writer)" + + - it: reader under Entra auth omits DATABASE_PASSWORD_READ_REPLICA + template: gateway/deployment.yaml + set: + database.writer.useAzureEntraAuth: true + database.reader.host: reader.example.com + database.reader.dbname: litellm + database.reader.useAzureEntraAuth: true + asserts: + - contains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_HOST_READ_REPLICA + value: reader.example.com + any: true + - contains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_USER_READ_REPLICA + valueFrom: + secretKeyRef: + name: litellm-reader-secret + key: username + any: true + - notContains: + path: spec.template.spec.containers[0].env + content: + name: DATABASE_PASSWORD_READ_REPLICA + any: true diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index 3f8aacfce17..998d225a317 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -145,6 +145,8 @@ database: dbname: "" schema: "" useIAMAuth: false + # Azure Database for PostgreSQL with a Microsoft Entra ID token; mutually exclusive with useIAMAuth + useAzureEntraAuth: false passwordSecret: name: litellm-writer-secret usernameKey: username @@ -159,6 +161,8 @@ database: dbname: "" schema: "" useIAMAuth: false + # Azure Database for PostgreSQL with a Microsoft Entra ID token; mutually exclusive with useIAMAuth + useAzureEntraAuth: false passwordSecret: name: litellm-reader-secret usernameKey: username diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260819000000_shadow_eval_max_budget/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260819000000_shadow_eval_max_budget/migration.sql new file mode 100644 index 00000000000..7b60dca9415 --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260819000000_shadow_eval_max_budget/migration.sql @@ -0,0 +1,5 @@ +-- AlterTable +ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN "max_budget" DOUBLE PRECISION; + +-- AlterTable +ALTER TABLE "LiteLLM_ShadowEvalAttempt" ADD COLUMN "shadow_cost" DOUBLE PRECISION NOT NULL DEFAULT 0; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 60058c777ca..d9959677116 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1502,7 +1502,8 @@ model LiteLLM_ShadowEvalJob { baseline_model String? // reverse only: the fixed model the router is judged against judge_model String shadow_percentage Float - max_turns Int // this key's sample budget: judge at most this many turns + max_turns Int // sample-count ceiling: the whole budget on pre-max_budget jobs, the error-loop valve otherwise + max_budget Float? // per-key USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets created_at DateTime @default(now()) created_by String? ends_at DateTime @@ -1525,6 +1526,7 @@ model LiteLLM_ShadowEvalAttempt { shadow_model String? confidence Float? judge_cost Float @default(0) + shadow_cost Float @default(0) error String? created_at DateTime @default(now()) diff --git a/litellm-proxy-extras/pyproject.toml b/litellm-proxy-extras/pyproject.toml index e1d62b70c29..26d42a33b29 100644 --- a/litellm-proxy-extras/pyproject.toml +++ b/litellm-proxy-extras/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "litellm-proxy-extras" -version = "0.4.87" +version = "0.4.88" description = "Additional files for the LiteLLM Proxy. Reduces the size of the main litellm package." readme = "README.md" requires-python = ">=3.9" @@ -26,7 +26,7 @@ required-version = ">=0.10.9" module-root = "" [tool.commitizen] -version = "0.4.87" +version = "0.4.88" version_files = [ "pyproject.toml:^version", "../pyproject.toml:litellm-proxy-extras==", diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs index 8b6896f3846..0c9faeda6e7 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -295,6 +295,8 @@ fn core_error_kind(error: &CoreError) -> &'static str { CoreError::Http { .. } => "HttpError", CoreError::InvalidResponse(_) => "InvalidResponse", CoreError::Network(_) => "NetworkError", + CoreError::Connect(_) => "ConnectError", CoreError::Routing(_) => "RoutingError", + CoreError::Unsupported(_) => "UnsupportedRequest", } } diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs index ffe2e0122c0..95df566dc53 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs @@ -324,6 +324,8 @@ fn core_error_kind(error: &CoreError) -> &'static str { CoreError::Http { .. } => "HttpError", CoreError::InvalidResponse(_) => "InvalidResponse", CoreError::Network(_) => "NetworkError", + CoreError::Connect(_) => "ConnectError", CoreError::Routing(_) => "RoutingError", + CoreError::Unsupported(_) => "UnsupportedRequest", } } diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index a34b2edd7b8..7e38d10c6ff 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -105,12 +105,20 @@ impl IntoResponse for MessagesRouteError { ), CoreError::Http { .. } | CoreError::Network(_) + | CoreError::Connect(_) | CoreError::InvalidResponse(_) | CoreError::InvalidType { .. } | CoreError::MissingField(_) => ( StatusCode::BAD_GATEWAY, "messages provider request failed".to_string(), ), + // The gateway has no Python implementation to decline to, so a + // request the core cannot serve is reported to the caller. The + // reason is a fixed internal string, never provider content. + CoreError::Unsupported(reason) => ( + StatusCode::BAD_REQUEST, + format!("messages request is not supported: {reason}"), + ), }; ( status, diff --git a/litellm-rust/crates/core/src/chat_completions/client.rs b/litellm-rust/crates/core/src/chat_completions/client.rs new file mode 100644 index 00000000000..f2ef73ed030 --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/client.rs @@ -0,0 +1,15 @@ +use std::sync::OnceLock; +use std::time::Duration; + +use crate::constants::{CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS, CHAT_COMPLETIONS_TIMEOUT_SECS}; + +pub(super) fn http_client() -> &'static reqwest::Client { + static CLIENT: OnceLock = OnceLock::new(); + CLIENT.get_or_init(|| { + reqwest::Client::builder() + .timeout(Duration::from_secs(CHAT_COMPLETIONS_TIMEOUT_SECS)) + .connect_timeout(Duration::from_secs(CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS)) + .build() + .unwrap_or_else(|_| reqwest::Client::new()) + }) +} diff --git a/litellm-rust/crates/core/src/chat_completions/common_utils.rs b/litellm-rust/crates/core/src/chat_completions/common_utils.rs new file mode 100644 index 00000000000..36eaf242a5a --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/common_utils.rs @@ -0,0 +1,28 @@ +use serde_json::{Map, Value}; + +use crate::error::CoreResult; +use crate::http_utils::string_headers as shared_string_headers; +use crate::providers::anthropic::chat_completions::transformation::ANTHROPIC_CHAT_COMPLETIONS_CONFIG; + +use super::transformation::ChatCompletionsProviderConfig; + +const HEADER_CONTEXT: &str = "chat completions"; + +pub(super) fn chat_completions_provider_config( + provider: &str, +) -> Option<&'static dyn ChatCompletionsProviderConfig> { + match provider { + "anthropic" => Some(&ANTHROPIC_CHAT_COMPLETIONS_CONFIG), + #[cfg(feature = "bedrock-auth")] + "bedrock" => Some( + &crate::providers::bedrock::chat_completions::transformation::BEDROCK_CHAT_COMPLETIONS_CONFIG, + ), + _ => None, + } +} + +pub(super) fn string_headers( + extra_headers: Option>, +) -> CoreResult> { + shared_string_headers(HEADER_CONTEXT, extra_headers) +} diff --git a/litellm-rust/crates/core/src/chat_completions/conversation.rs b/litellm-rust/crates/core/src/chat_completions/conversation.rs new file mode 100644 index 00000000000..f7bdc60af37 --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/conversation.rs @@ -0,0 +1,254 @@ +//! Provider-neutral conversation shape. +//! +//! Both Anthropic Messages and Bedrock Converse want the same thing out of an +//! OpenAI message list: the system prompt lifted out, consecutive same-role +//! turns merged, and text blocks that are never empty. That normalization is +//! shared here so a provider config only renders the result into its own wire +//! shape. +//! +//! Mirrors Python's `anthropic_messages_pt` / +//! `_bedrock_converse_messages_pt` for the text-only surface this route +//! accepts; anything richer is declined upstream by the capability gate. + +use crate::constants::EMPTY_TEXT_PLACEHOLDER; + +use super::types::{ChatMessage, ChatMessageContent}; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum TurnRole { + User, + Assistant, +} + +impl TurnRole { + pub fn as_str(self) -> &'static str { + match self { + Self::User => "user", + Self::Assistant => "assistant", + } + } +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub struct Turn { + pub role: TurnRole, + pub texts: Vec, +} + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub struct Conversation { + pub system: Vec, + pub turns: Vec, +} + +/// True when the conversation can be sent as-is. +/// +/// Python inserts a placeholder first user turn only under +/// `litellm.modify_params`, which the core cannot see, so a conversation that +/// does not open on a user turn is declined rather than guessed at. +impl Conversation { + pub fn opens_on_user_turn(&self) -> bool { + self.turns + .first() + .is_some_and(|turn| turn.role == TurnRole::User) + } +} + +fn message_texts(content: &ChatMessageContent) -> Vec { + match content { + ChatMessageContent::Text(text) => vec![text.clone()], + ChatMessageContent::Parts(parts) => parts + .iter() + .filter_map(|part| part.get("text").and_then(|text| text.as_str())) + .map(str::to_string) + .collect(), + } +} + +/// Python rewrites empty or whitespace-only text rather than dropping it, so an +/// entirely empty content list never reaches a provider that rejects one. +fn sanitize(text: String) -> String { + if text.trim().is_empty() { + return EMPTY_TEXT_PLACEHOLDER.to_string(); + } + text +} + +pub fn build_conversation(messages: &[ChatMessage]) -> Conversation { + let system = messages + .iter() + .filter(|message| message.role == "system") + .filter_map(|message| message.content.as_ref()) + .flat_map(message_texts) + .filter(|text| !text.is_empty()) + .collect(); + + let turns = messages + .iter() + .filter(|message| message.role != "system") + .fold(Vec::::new(), |mut turns, message| { + let role = if message.role == "assistant" { + TurnRole::Assistant + } else { + TurnRole::User + }; + let texts = message + .content + .as_ref() + .map(message_texts) + .unwrap_or_default() + .into_iter() + .map(sanitize); + match turns.last_mut() { + Some(last) if last.role == role => last.texts.extend(texts), + _ => turns.push(Turn { + role, + texts: texts.collect(), + }), + } + turns + }); + + // Anthropic and Bedrock both reject trailing whitespace on the final + // assistant turn, so Python right-strips it there; mirror that exactly. + let turns = match turns.split_last() { + Some((last, rest)) if last.role == TurnRole::Assistant => rest + .iter() + .cloned() + .chain([Turn { + role: last.role, + texts: last + .texts + .iter() + .map(|text| text.trim_end().to_string()) + .collect(), + }]) + .collect(), + _ => turns, + }; + + Conversation { system, turns } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + fn messages(value: serde_json::Value) -> Vec { + serde_json::from_value(value).expect("valid messages") + } + + #[test] + fn lifts_system_messages_out_of_the_turn_list() { + let conversation = build_conversation(&messages(json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"} + ]))); + assert_eq!(conversation.system, vec!["be terse".to_string()]); + assert_eq!( + conversation.turns, + vec![Turn { + role: TurnRole::User, + texts: vec!["hi".to_string()] + }] + ); + } + + #[test] + fn merges_consecutive_same_role_turns() { + let conversation = build_conversation(&messages(json!([ + {"role": "user", "content": "one"}, + {"role": "user", "content": "two"}, + {"role": "assistant", "content": "ack"}, + {"role": "user", "content": "three"} + ]))); + assert_eq!( + conversation.turns, + vec![ + Turn { + role: TurnRole::User, + texts: vec!["one".to_string(), "two".to_string()] + }, + Turn { + role: TurnRole::Assistant, + texts: vec!["ack".to_string()] + }, + Turn { + role: TurnRole::User, + texts: vec!["three".to_string()] + }, + ] + ); + } + + #[test] + fn flattens_text_parts_in_order() { + let conversation = build_conversation(&messages(json!([ + {"role": "user", "content": [ + {"type": "text", "text": "first"}, + {"type": "text", "text": "second"} + ]} + ]))); + assert_eq!( + conversation.turns[0].texts, + vec!["first".to_string(), "second".to_string()] + ); + } + + #[test] + fn rewrites_empty_and_whitespace_only_text_to_the_python_placeholder() { + let conversation = build_conversation(&messages(json!([ + {"role": "user", "content": ""}, + {"role": "assistant", "content": " "}, + {"role": "user", "content": "real"} + ]))); + assert_eq!(conversation.turns[0].texts, vec![EMPTY_TEXT_PLACEHOLDER]); + assert_eq!(conversation.turns[1].texts, vec![EMPTY_TEXT_PLACEHOLDER]); + } + + #[test] + fn right_strips_only_the_final_assistant_turn() { + let conversation = build_conversation(&messages(json!([ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "kept "}, + {"role": "user", "content": "more"}, + {"role": "assistant", "content": "stripped "} + ]))); + assert_eq!(conversation.turns[1].texts, vec!["kept ".to_string()]); + assert_eq!(conversation.turns[3].texts, vec!["stripped".to_string()]); + } + + #[test] + fn does_not_strip_when_the_last_turn_is_a_user_turn() { + let conversation = build_conversation(&messages(json!([ + {"role": "assistant", "content": "kept "}, + {"role": "user", "content": "hi "} + ]))); + assert_eq!(conversation.turns[0].texts, vec!["kept ".to_string()]); + assert_eq!(conversation.turns[1].texts, vec!["hi ".to_string()]); + } + + #[test] + fn reports_whether_the_conversation_opens_on_a_user_turn() { + assert!( + build_conversation(&messages(json!([{"role": "user", "content": "hi"}]))) + .opens_on_user_turn() + ); + assert!( + !build_conversation(&messages(json!([{"role": "assistant", "content": "hi"}]))) + .opens_on_user_turn() + ); + assert!(!Conversation::default().opens_on_user_turn()); + } + + #[test] + fn drops_empty_system_text_the_way_python_skips_empty_system_blocks() { + let conversation = build_conversation(&messages(json!([ + {"role": "system", "content": ""}, + {"role": "system", "content": "kept"}, + {"role": "user", "content": "hi"} + ]))); + assert_eq!(conversation.system, vec!["kept".to_string()]); + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/handler.rs b/litellm-rust/crates/core/src/chat_completions/handler.rs new file mode 100644 index 00000000000..afc4529fd26 --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/handler.rs @@ -0,0 +1,147 @@ +use serde_json::Value; + +use crate::error::{CoreError, CoreResult}; +use crate::http_utils::truncate_error_body; + +use super::client::http_client; +use super::transformation::ChatCompletionsAuth; +use super::types::{ + ChatCompletionsResponse, ProviderChatCompletionsRequest, ProviderChatResponseData, +}; + +pub(super) async fn execute_chat_completions_provider_call( + request: ProviderChatCompletionsRequest, +) -> CoreResult { + let body = serde_json::to_vec(&request.body).map_err(|err| { + CoreError::InvalidRequest(format!( + "failed to serialize chat completions request: {err}" + )) + })?; + let headers = signed_headers(&request, &body).await?; + + let mut request_builder = http_client().post(&request.url).body(body); + for (key, value) in &headers { + request_builder = request_builder.header(key, value); + } + if let Some(duration) = request.timeout { + request_builder = request_builder.timeout(duration); + } + + let response = request_builder.send().await.map_err(|err| { + // Failing to establish the connection means the request never went out, + // so the host can still serve it. Everything else here, a timeout + // above all, may have reached the provider and been answered. + if err.is_connect() || err.is_builder() { + CoreError::Connect(err.to_string()) + } else { + CoreError::Network(err.to_string()) + } + })?; + + let status = response.status(); + let text = response + .text() + .await + .map_err(|err| CoreError::Network(err.to_string()))?; + + if !status.is_success() { + return Err(CoreError::Http { + status: status.as_u16(), + body: truncate_error_body(&text), + }); + } + + let body: Value = serde_json::from_str(&text).map_err(|err| { + CoreError::InvalidResponse(format!("invalid chat completions response JSON: {err}")) + })?; + request + .config + .transform_response(&request.model, ProviderChatResponseData { body }) + .map_err(as_response_error) +} + +/// Re-tag an error raised while normalizing a response the provider already +/// returned. +/// +/// A config reports the same variants on either side of the call: a missing +/// field or an unsupported block can mean "this request cannot be translated" +/// during prepare and "this response cannot be normalized" here. Only the +/// second kind has already been billed, and a host that keeps a reference +/// implementation must not retry those, so collapse them to one variant that +/// can only mean the provider was already called. +pub(super) fn as_response_error(err: CoreError) -> CoreError { + match err { + already @ (CoreError::InvalidResponse(_) | CoreError::Http { .. }) => already, + other => CoreError::InvalidResponse(other.to_string()), + } +} + +#[cfg(feature = "bedrock-auth")] +pub(super) async fn signed_headers( + request: &ProviderChatCompletionsRequest, + body: &[u8], +) -> CoreResult> { + use std::collections::BTreeMap; + use std::time::SystemTime; + + use crate::providers::bedrock::aws_base::{ + aws_auth_config, aws_signature_headers, host_supplied_credentials, + is_sigv4_computed_header, resolve_credentials, sign_bedrock_post, + }; + + let ChatCompletionsAuth::AwsSigV4 { region } = &request.auth else { + return Ok(request.upstream_headers.clone()); + }; + // Reattaching a header the signer also emits would put both copies on the + // wire, and Bedrock rejects that pair. Python instead drops the caller's + // copy and prefers a forwarded Authorization over the signature, so leave + // the request to Python rather than serving it a different way here. + if request + .upstream_headers + .iter() + .any(|(name, _)| is_sigv4_computed_header(name)) + { + return Err(CoreError::Unsupported( + "request forwards a header AWS SigV4 computes", + )); + } + let env_lookup = |key: &str| std::env::var(key).ok(); + let unsigned: BTreeMap = request.upstream_headers.iter().cloned().collect(); + // A host with its own resolution chain hands the result down; only fall + // back to deriving credentials here when it supplied none. + let credentials = match host_supplied_credentials(&request.optional_params) { + Some(credentials) => credentials, + None => { + resolve_credentials( + aws_auth_config(&request.optional_params, &env_lookup), + &env_lookup, + ) + .await? + } + }; + let signature = sign_bedrock_post( + &request.url, + body, + &aws_signature_headers(&unsigned), + region, + &credentials, + SystemTime::now(), + )?; + // Every original header goes back on the wire alongside the computed ones, + // as Python reattaches them. The guard above already rejected the names + // that would collide, so no name appears twice. + Ok(unsigned.into_iter().chain(signature).collect()) +} + +#[cfg(not(feature = "bedrock-auth"))] +pub(super) async fn signed_headers( + request: &ProviderChatCompletionsRequest, + _body: &[u8], +) -> CoreResult> { + match &request.auth { + ChatCompletionsAuth::AwsSigV4 { .. } => Err(CoreError::Unsupported( + "AWS SigV4 requires the bedrock-auth feature", + )), + _ => Ok(request.upstream_headers.clone()), + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/mod.rs b/litellm-rust/crates/core/src/chat_completions/mod.rs new file mode 100644 index 00000000000..f30ac1a24bf --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/mod.rs @@ -0,0 +1,59 @@ +//! The `/chat/completions` call, the Rust equivalent of Python's +//! `litellm.completion()`. +//! +//! [`chat_completions`] is the top-level entrypoint: give it a model, the +//! OpenAI-shaped message list, the provider-mapped optional params, and +//! credentials, and it resolves the provider, translates the conversation, +//! calls the provider, and returns a typed OpenAI-shaped response. + +mod client; +mod common_utils; +pub mod conversation; +pub(crate) mod handler; +mod prepare; +pub mod response_utils; +pub mod transformation; +pub mod types; + +use serde_json::{Map, Value}; + +use crate::error::CoreResult; + +use handler::execute_chat_completions_provider_call; +use prepare::{parse_messages, prepare_chat_completions_call, resolve_provider_config}; +use types::{ChatCompletionsRequest, ChatCompletionsResponse}; + +pub async fn chat_completions( + request: ChatCompletionsRequest<'_>, +) -> CoreResult { + execute_chat_completions_provider_call(prepare_chat_completions_call(request)?).await +} + +/// Whether the core would accept this request, without resolving credentials or +/// touching the network. +/// +/// A host that keeps the Python implementation asks this first so it can emit +/// its pre-call logging exactly once, on whichever path is about to run. +/// Returns the decline reason, or `None` when the request is accepted. +pub fn chat_completions_decline_reason( + model: &str, + custom_llm_provider: Option<&str>, + messages: Value, + optional_params: &Map, +) -> Option<&'static str> { + let Ok((_, config)) = resolve_provider_config(model, custom_llm_provider) else { + return Some("provider is not on the rust chat completions path"); + }; + let Ok(messages) = parse_messages(messages) else { + return Some("unreadable message list"); + }; + if messages.is_empty() { + return Some("empty message list"); + } + config + .unsupported_reason(&messages, optional_params) + .map(|reason| reason.0) +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/core/src/chat_completions/prepare.rs b/litellm-rust/crates/core/src/chat_completions/prepare.rs new file mode 100644 index 00000000000..1e1c8d1bafd --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/prepare.rs @@ -0,0 +1,118 @@ +use serde_json::Value; + +use crate::error::{CoreError, CoreResult}; +use crate::http_utils::has_header; +use crate::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; + +use super::common_utils::{chat_completions_provider_config, string_headers}; +use super::transformation::{ChatCompletionsAuth, ChatCompletionsProviderConfig}; +use super::types::{ChatCompletionsRequest, ChatMessage, ProviderChatCompletionsRequest}; + +pub(super) fn resolve_provider_config<'a>( + model: &'a str, + custom_llm_provider: Option<&'a str>, +) -> CoreResult<(String, &'static dyn ChatCompletionsProviderConfig)> { + let provider_info = get_custom_llm_provider(model, custom_llm_provider) + .or_else(|| { + custom_llm_provider.map(|provider| CustomLlmProvider { + model, + custom_llm_provider: provider, + }) + }) + .ok_or_else(|| { + CoreError::InvalidProvider( + "unable to resolve custom_llm_provider for chat completions request".to_string(), + ) + })?; + let config = chat_completions_provider_config(provider_info.custom_llm_provider) + .ok_or_else(|| CoreError::InvalidProvider(provider_info.custom_llm_provider.to_string()))?; + Ok((provider_info.model.to_string(), config)) +} + +pub(super) fn parse_messages(messages: Value) -> CoreResult> { + serde_json::from_value(messages).map_err(|err| { + CoreError::InvalidRequest(format!("invalid chat completions messages: {err}")) + }) +} + +pub(super) fn prepare_chat_completions_call( + request: ChatCompletionsRequest<'_>, +) -> CoreResult { + let (model, config) = resolve_provider_config(request.model, request.custom_llm_provider)?; + let env_lookup = |key: &str| std::env::var(key).ok(); + + let messages = parse_messages(request.messages)?; + if messages.is_empty() { + return Err(CoreError::InvalidRequest( + "chat completions requires at least one message".to_string(), + )); + } + if let Some(reason) = config.unsupported_reason(&messages, &request.optional_params) { + return Err(CoreError::Unsupported(reason.0)); + } + + let mut headers = string_headers(request.extra_headers)?; + let auth = config.auth( + request.api_key, + &model, + &request.optional_params, + &env_lookup, + )?; + match &auth { + ChatCompletionsAuth::Header { name, value } => { + // The deployment's credential replaces whatever the caller forwarded + // under the same name, mirroring Python's + // `{**headers, **anthropic_headers}`: letting a request header win + // would let its sender choose the principal the call bills to. + // + // The exception is a scheme the provider hands off to entirely, such + // as an Anthropic OAuth bearer, where Python drops `x-api-key` + // instead of resolving one. Re-adding it there would put the + // credential into a header the host removed on purpose. + if !config.defers_to_forwarded_auth(&headers) { + headers.retain(|(header, _)| !header.eq_ignore_ascii_case(name)); + headers.push(((*name).to_string(), value.clone())); + } + } + ChatCompletionsAuth::Bearer { token } => { + // Bedrock's `get_request_headers` assigns `headers["Authorization"]` + // unconditionally once a bearer token resolves, so the deployment's + // identity outranks whatever the caller forwarded. Keeping the + // caller's would bill and authorize the call as a different + // principal than the same deployment uses on Python. + // + // The `Header` arm below keeps the opposite precedence on purpose: + // Anthropic's transform honours a forwarded OAuth bearer. + headers.retain(|(name, _)| !name.eq_ignore_ascii_case("authorization")); + headers.push(("authorization".to_string(), format!("Bearer {token}"))); + } + // SigV4 signs the serialized body, so the handler adds its headers. + ChatCompletionsAuth::AwsSigV4 { .. } => {} + } + + for (name, value) in config.default_headers() { + if !has_header(&headers, name) { + headers.push(((*name).to_string(), (*value).to_string())); + } + } + + let url = config.complete_url( + request.api_base, + &model, + &request.optional_params, + &env_lookup, + )?; + let transformed = + config.transform_request(&model, messages, request.optional_params.clone())?; + + Ok(ProviderChatCompletionsRequest { + model, + config, + url, + body: transformed.body, + upstream_headers: headers, + auth, + optional_params: request.optional_params, + timeout: request.timeout, + }) +} diff --git a/litellm-rust/crates/core/src/chat_completions/response_utils.rs b/litellm-rust/crates/core/src/chat_completions/response_utils.rs new file mode 100644 index 00000000000..1ada5d43980 --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/response_utils.rs @@ -0,0 +1,101 @@ +//! Response normalization shared by every chat completions provider config. + +use std::time::{SystemTime, UNIX_EPOCH}; + +use super::types::{ChatCompletionsUsage, PromptTokensDetails}; + +/// OpenAI finish reasons, mirroring Python's `_FINISH_REASON_MAP` for the +/// reasons the providers on this route can emit. Python warns and falls back to +/// `stop` for anything unmapped, so do the same. +const FINISH_REASONS: &[(&str, &str)] = &[ + ("end_turn", "stop"), + ("stop_sequence", "stop"), + ("max_tokens", "length"), + ("refusal", "content_filter"), + ("compaction", "length"), + ("guardrail_intervened", "content_filter"), + ("content_filtered", "content_filter"), + ("content_filter", "content_filter"), + ("stop", "stop"), + ("length", "length"), +]; + +pub fn finish_reason_for(provider_reason: &str) -> &'static str { + FINISH_REASONS + .iter() + .find(|(reason, _)| *reason == provider_reason) + .map_or("stop", |(_, mapped)| *mapped) +} + +/// Python folds cache tokens into `prompt_tokens` and reports the split under +/// `prompt_tokens_details`; mirror that so cost tracking agrees on both paths. +pub fn usage_from_parts( + input_tokens: u64, + output_tokens: u64, + cache_read_tokens: u64, + cache_creation_tokens: u64, +) -> ChatCompletionsUsage { + let prompt_tokens = input_tokens + cache_read_tokens + cache_creation_tokens; + ChatCompletionsUsage { + prompt_tokens, + completion_tokens: output_tokens, + total_tokens: prompt_tokens + output_tokens, + prompt_tokens_details: PromptTokensDetails { + cached_tokens: cache_read_tokens, + cache_creation_tokens, + text_tokens: input_tokens, + }, + } +} + +pub fn unix_now() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map_or(0, |elapsed| elapsed.as_secs()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn maps_every_reason_the_route_can_observe() { + assert_eq!(finish_reason_for("end_turn"), "stop"); + assert_eq!(finish_reason_for("stop_sequence"), "stop"); + assert_eq!(finish_reason_for("max_tokens"), "length"); + assert_eq!(finish_reason_for("refusal"), "content_filter"); + assert_eq!(finish_reason_for("guardrail_intervened"), "content_filter"); + // Converse emits these two, and folding them into `stop` would report a + // filtered completion as a normal one. + assert_eq!(finish_reason_for("content_filtered"), "content_filter"); + assert_eq!(finish_reason_for("content_filter"), "content_filter"); + } + + #[test] + fn defaults_an_unmapped_reason_to_stop_like_python() { + // Python warns and falls back to `stop` for a reason its own map does + // not carry, so only a reason absent from `_FINISH_REASON_MAP` belongs + // here. + assert_eq!(finish_reason_for("something_new"), "stop"); + assert_eq!(finish_reason_for(""), "stop"); + } + + #[test] + fn folds_cache_tokens_into_prompt_tokens() { + let usage = usage_from_parts(10, 4, 7, 3); + assert_eq!(usage.prompt_tokens, 20); + assert_eq!(usage.completion_tokens, 4); + assert_eq!(usage.total_tokens, 24); + assert_eq!(usage.prompt_tokens_details.cached_tokens, 7); + assert_eq!(usage.prompt_tokens_details.cache_creation_tokens, 3); + assert_eq!(usage.prompt_tokens_details.text_tokens, 10); + } + + #[test] + fn reports_raw_input_tokens_when_no_cache_is_involved() { + let usage = usage_from_parts(12, 5, 0, 0); + assert_eq!(usage.prompt_tokens, 12); + assert_eq!(usage.total_tokens, 17); + assert_eq!(usage.prompt_tokens_details.text_tokens, 12); + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/tests.rs b/litellm-rust/crates/core/src/chat_completions/tests.rs new file mode 100644 index 00000000000..e2383723cb0 --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/tests.rs @@ -0,0 +1,820 @@ +use serde_json::{Map, Value, json}; + +use crate::error::CoreError; + +use super::prepare::prepare_chat_completions_call; +use super::transformation::ChatCompletionsAuth; +use super::types::ChatCompletionsRequest; + +fn request<'a>( + model: &'a str, + provider: Option<&'a str>, + messages: Value, + optional_params: Value, +) -> ChatCompletionsRequest<'a> { + ChatCompletionsRequest { + model, + messages, + optional_params: match optional_params { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + }, + api_key: Some("sk-test"), + api_base: None, + custom_llm_provider: provider, + extra_headers: None, + timeout: None, + } +} + +/// `ProviderChatCompletionsRequest` deliberately has no `Debug` (its headers +/// carry resolved credentials), so unwrap the failure case by hand. +fn decline(request: ChatCompletionsRequest<'_>) -> CoreError { + match prepare_chat_completions_call(request) { + Err(error) => error, + Ok(prepared) => panic!("expected a decline, prepared a call to {}", prepared.url), + } +} + +#[test] +fn resolves_the_provider_from_the_model_prefix() { + let prepared = prepare_chat_completions_call(request( + "anthropic/claude-sonnet-4-5", + None, + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + )) + .expect("prepares"); + assert_eq!(prepared.model, "claude-sonnet-4-5"); + assert_eq!(prepared.url, "https://api.anthropic.com/v1/messages"); + assert_eq!(prepared.body["model"], json!("claude-sonnet-4-5")); +} + +#[test] +fn strips_an_explicit_provider_prefix_from_the_model() { + let prepared = prepare_chat_completions_call(request( + "anthropic/claude-sonnet-4-5", + Some("anthropic"), + json!([{"role": "user", "content": "hi"}]), + json!({}), + )) + .expect("prepares"); + assert_eq!(prepared.model, "claude-sonnet-4-5"); +} + +#[test] +fn adds_the_auth_and_default_headers() { + let prepared = prepare_chat_completions_call(request( + "claude-sonnet-4-5", + Some("anthropic"), + json!([{"role": "user", "content": "hi"}]), + json!({}), + )) + .expect("prepares"); + assert!( + prepared + .upstream_headers + .contains(&("x-api-key".to_string(), "sk-test".to_string())) + ); + assert!( + prepared + .upstream_headers + .contains(&("anthropic-version".to_string(), "2023-06-01".to_string())) + ); + assert!(matches!( + prepared.auth, + ChatCompletionsAuth::Header { + name: "x-api-key", + .. + } + )); +} + +#[test] +fn the_deployment_credential_replaces_a_caller_supplied_auth_header() { + // Python builds `{**headers, **anthropic_headers}`, so the deployment's key + // overwrites a forwarded one. Honouring the caller's would let whoever sends + // the request choose the Anthropic principal it bills to. + let mut call = request( + "claude-sonnet-4-5", + Some("anthropic"), + json!([{"role": "user", "content": "hi"}]), + json!({}), + ); + call.extra_headers = Some(Map::from_iter([( + "X-Api-Key".to_string(), + json!("sk-caller"), + )])); + let prepared = prepare_chat_completions_call(call).expect("prepares"); + let keys: Vec<_> = prepared + .upstream_headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) + .collect(); + assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers); + assert_eq!(keys[0].1, "sk-test"); +} + +#[test] +fn a_forwarded_authorization_header_suppresses_the_resolved_api_key_header() { + // Anthropic's `validate_environment` pops `x-api-key` and sets `authorization` + // for an OAuth token, so re-adding the key here would put the credential into + // a header the host removed on purpose. + let mut call = request( + "claude-sonnet-4-5", + Some("anthropic"), + json!([{"role": "user", "content": "hi"}]), + json!({}), + ); + call.extra_headers = Some(Map::from_iter([ + ( + "Authorization".to_string(), + json!("Bearer sk-ant-oat01-token"), + ), + ("X-Api-Key".to_string(), json!("sk-caller")), + ])); + let prepared = prepare_chat_completions_call(call).expect("prepares"); + assert!( + !prepared + .upstream_headers + .iter() + .any(|(name, value)| name.eq_ignore_ascii_case("x-api-key") && value == "sk-test"), + "the resolved key must not be applied over an OAuth bearer, got {:?}", + prepared.upstream_headers + ); + assert!( + prepared + .upstream_headers + .iter() + .any(|(name, value)| name.eq_ignore_ascii_case("authorization") + && value == "Bearer sk-ant-oat01-token") + ); +} + +#[test] +fn an_unrelated_forwarded_authorization_does_not_defer_the_resolved_key() { + // Only an OAuth bearer replaces the credential. Python sends the deployment's + // `x-api-key` alongside any other forwarded `authorization`, so deferring on + // the mere presence of that header would drop the deployment's auth. + let mut call = request( + "claude-sonnet-4-5", + Some("anthropic"), + json!([{"role": "user", "content": "hi"}]), + json!({}), + ); + call.extra_headers = Some(Map::from_iter([ + ("Authorization".to_string(), json!("Bearer unrelated")), + ("X-Api-Key".to_string(), json!("sk-caller")), + ])); + let prepared = prepare_chat_completions_call(call).expect("prepares"); + let keys: Vec<_> = prepared + .upstream_headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) + .collect(); + assert_eq!(keys.len(), 1, "got {:?}", prepared.upstream_headers); + assert_eq!(keys[0].1, "sk-test"); + assert!( + prepared + .upstream_headers + .iter() + .any(|(name, value)| name.eq_ignore_ascii_case("authorization") + && value == "Bearer unrelated"), + "the unrelated authorization must survive, got {:?}", + prepared.upstream_headers + ); +} + +#[test] +fn declines_an_unsupported_request_before_resolving_credentials() { + let mut call = request( + "claude-sonnet-4-5", + Some("anthropic"), + json!([{"role": "user", "content": "hi"}]), + json!({"stream": true}), + ); + call.api_key = None; + // No api_key is set and no env is consulted: the gate must run first, so the + // error is the decline rather than a missing-credential error. + assert_eq!(decline(call), CoreError::Unsupported("streaming")); +} + +#[test] +fn rejects_an_unknown_provider() { + assert_eq!( + decline(request( + "openai/gpt-4o", + None, + json!([{"role": "user", "content": "hi"}]), + json!({}), + )), + CoreError::InvalidProvider("openai".to_string()) + ); +} + +#[test] +fn rejects_a_model_with_no_resolvable_provider() { + assert!(matches!( + decline(request( + "claude-sonnet-4-5", + None, + json!([{"role": "user", "content": "hi"}]), + json!({}), + )), + CoreError::InvalidProvider(_) + )); +} + +#[test] +fn rejects_an_empty_or_malformed_message_list() { + assert_eq!( + decline(request( + "anthropic/claude-sonnet-4-5", + None, + json!([]), + json!({}), + )), + CoreError::InvalidRequest("chat completions requires at least one message".to_string()) + ); + assert!(matches!( + decline(request( + "anthropic/claude-sonnet-4-5", + None, + json!("not a list"), + json!({}), + )), + CoreError::InvalidRequest(_) + )); +} + +#[test] +fn rejects_non_string_extra_headers() { + let mut call = request( + "anthropic/claude-sonnet-4-5", + None, + json!([{"role": "user", "content": "hi"}]), + json!({}), + ); + call.extra_headers = Some(Map::from_iter([("x-trace".to_string(), json!(7))])); + assert_eq!( + decline(call), + CoreError::InvalidRequest( + "chat completions extra_headers.x-trace must be a string, got number".to_string() + ) + ); +} + +#[cfg(feature = "bedrock-auth")] +#[test] +fn prepares_a_bedrock_call_without_resolving_credentials() { + let mut call = request( + "bedrock/us-east-1/anthropic.claude-v2", + None, + json!([{"role": "user", "content": "hi"}]), + json!({"maxTokens": 16}), + ); + call.api_key = None; + let prepared = prepare_chat_completions_call(call).expect("prepares"); + assert_eq!( + prepared.url, + "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse" + ); + assert_eq!( + prepared.auth, + ChatCompletionsAuth::AwsSigV4 { + region: "us-east-1".to_string() + } + ); + // SigV4 signs the serialized body, so prepare must not have added an + // Authorization header; the handler does it. + assert!( + !prepared + .upstream_headers + .iter() + .any(|(name, _)| name.eq_ignore_ascii_case("authorization")) + ); + assert_eq!(prepared.body["inferenceConfig"], json!({"maxTokens": 16})); +} + +#[cfg(feature = "bedrock-auth")] +#[tokio::test] +async fn a_forwarded_client_header_does_not_enter_the_bedrock_signature() { + // Python signs only the AWS header set and reattaches the rest, so a header + // the caller forwarded rides along without joining the canonical request. + // Signing it makes Converse 403 on a deployment that works on Python. + let mut call = request( + "bedrock/us-east-1/anthropic.claude-v2", + None, + json!([{"role": "user", "content": "hi"}]), + json!({ + "maxTokens": 16, + "aws_access_key_id": "AKIDEXAMPLE", + "aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY" + }), + ); + // A key would resolve to a bearer token and never reach the signer. + call.api_key = None; + call.extra_headers = Some(Map::from_iter([( + "x-request-id".to_string(), + json!("abc-123"), + )])); + let prepared = prepare_chat_completions_call(call).expect("prepares"); + let signed = super::handler::signed_headers(&prepared, br#"{"a":1}"#) + .await + .expect("signs"); + + let authorization = signed + .iter() + .find(|(name, _)| name.eq_ignore_ascii_case("authorization")) + .map(|(_, value)| value.clone()) + .expect("carries an authorization header"); + assert!( + authorization.starts_with("AWS4-HMAC-SHA256"), + "expected a SigV4 signature, got {authorization}" + ); + assert!( + !authorization.contains("x-request-id"), + "forwarded header reached SignedHeaders: {authorization}" + ); + // It still goes on the wire, it is just not part of the signature. + assert!( + signed + .iter() + .any(|(name, value)| name == "x-request-id" && value == "abc-123"), + "forwarded header was dropped instead of reattached" + ); +} + +#[cfg(feature = "bedrock-auth")] +#[tokio::test] +async fn a_forwarded_header_the_signer_computes_declines_to_python() { + // Reattaching the caller's copy next to the computed one puts the name on + // the wire twice and Bedrock rejects the pair, so a request carrying one + // has to go to Python instead of being signed here. + for forwarded in [ + "Authorization", + "x-amz-date", + "x-amz-security-token", + "Date", + ] { + let mut call = request( + "bedrock/us-east-1/anthropic.claude-v2", + None, + json!([{"role": "user", "content": "hi"}]), + json!({ + "maxTokens": 16, + "aws_access_key_id": "AKIDEXAMPLE", + "aws_secret_access_key": "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY" + }), + ); + call.api_key = None; + call.extra_headers = Some(Map::from_iter([(forwarded.to_string(), json!("forged"))])); + let prepared = prepare_chat_completions_call(call).expect("prepares"); + let error = super::handler::signed_headers(&prepared, br#"{"a":1}"#) + .await + .expect_err("{forwarded} should decline instead of being signed"); + assert!( + matches!(error, CoreError::Unsupported(_)), + "{forwarded} declined as {error:?}, which the host would not fall back on" + ); + } +} + +#[cfg(feature = "bedrock-auth")] +#[test] +fn a_bedrock_deployment_bearer_outranks_a_forwarded_authorization() { + // `get_request_headers` assigns `headers["Authorization"]` unconditionally + // once a bearer token resolves, so the deployment's identity wins on + // Python. Keeping the caller's would authorize and bill the call as a + // different principal, and only when the deployment carries `rust: true`. + let mut call = request( + "bedrock/us-east-1/anthropic.claude-v2", + None, + json!([{"role": "user", "content": "hi"}]), + json!({"maxTokens": 16}), + ); + call.extra_headers = Some(Map::from_iter([( + "Authorization".to_string(), + json!("Bearer caller-supplied"), + )])); + let prepared = prepare_chat_completions_call(call).expect("prepares"); + let authorizations: Vec<_> = prepared + .upstream_headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case("authorization")) + .map(|(_, value)| value.as_str()) + .collect(); + assert_eq!( + authorizations, + vec!["Bearer sk-test"], + "the deployment token must be the only authorization on the wire" + ); +} + +#[test] +fn an_anthropic_forwarded_oauth_bearer_still_outranks_the_resolved_key() { + // The opposite precedence, and deliberate: Anthropic's own transform + // honours a forwarded OAuth bearer, so the Bedrock fix above must not be + // generalized into a rule that the configured key always wins. + // + // An OAuth bearer is the whole of that exception. This forwarded a plain + // `x-api-key` until round 17, which read as the same claim and was not: + // Python overwrites a forwarded `x-api-key` with the deployment's. + let mut call = request( + "claude-sonnet-4-5", + Some("anthropic"), + json!([{"role": "user", "content": "hi"}]), + json!({}), + ); + call.extra_headers = Some(Map::from_iter([( + "authorization".to_string(), + json!("Bearer sk-ant-oat01-forwarded"), + )])); + let prepared = prepare_chat_completions_call(call).expect("prepares"); + let keys: Vec<_> = prepared + .upstream_headers + .iter() + .filter(|(name, _)| name.eq_ignore_ascii_case("x-api-key")) + .map(|(_, value)| value.as_str()) + .collect(); + assert!(keys.is_empty(), "got {:?}", prepared.upstream_headers); + assert!( + prepared + .upstream_headers + .iter() + .any(|(name, value)| name.eq_ignore_ascii_case("authorization") + && value == "Bearer sk-ant-oat01-forwarded") + ); +} + +#[cfg(feature = "bedrock-auth")] +#[test] +fn a_bedrock_api_key_is_sent_as_a_bearer_token_instead_of_being_signed() { + // The configured bearer identity has its own account and quota boundary, + // so a request carrying one must not be signed as whatever principal the + // host's AWS credentials resolve to. + let prepared = prepare_chat_completions_call(request( + "bedrock/us-east-1/anthropic.claude-v2", + None, + json!([{"role": "user", "content": "hi"}]), + json!({"maxTokens": 16}), + )) + .expect("prepares"); + assert_eq!( + prepared.auth, + ChatCompletionsAuth::Bearer { + token: "sk-test".to_string() + } + ); + assert!( + prepared + .upstream_headers + .iter() + .any(|(name, value)| name.eq_ignore_ascii_case("authorization") + && value == "Bearer sk-test"), + "prepare did not carry the bearer token" + ); +} + +fn decline_reason( + model: &str, + provider: Option<&str>, + messages: Value, + params: Value, +) -> Option<&'static str> { + let params = match params { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + }; + super::chat_completions_decline_reason(model, provider, messages, ¶ms) +} + +#[test] +fn the_gate_accepts_what_prepare_accepts() { + assert_eq!( + decline_reason( + "anthropic/claude-sonnet-4-5", + None, + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + ), + None + ); +} + +#[test] +fn the_gate_declines_without_resolving_credentials_or_calling_out() { + assert_eq!( + decline_reason( + "anthropic/claude-sonnet-4-5", + None, + json!([{"role": "user", "content": "hi"}]), + json!({"stream": true}), + ), + Some("streaming") + ); + assert_eq!( + decline_reason( + "openai/gpt-4o", + None, + json!([{"role": "user", "content": "hi"}]), + json!({}), + ), + Some("provider is not on the rust chat completions path") + ); + assert_eq!( + decline_reason( + "claude-sonnet-4-5", + None, + json!([{"role": "user", "content": "hi"}]), + json!({}), + ), + Some("provider is not on the rust chat completions path") + ); + assert_eq!( + decline_reason( + "anthropic/claude-sonnet-4-5", + None, + json!("nope"), + json!({}) + ), + Some("unreadable message list") + ); + assert_eq!( + decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})), + Some("empty message list") + ); +} + +#[test] +fn the_gate_agrees_with_prepare_on_every_case_it_accepts() { + // A gate that accepts what prepare then declines would make the host emit + // its pre-call logging on a path that falls back, so pin the agreement. + for (messages, params) in [ + ( + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 8}), + ), + ( + json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]), + json!({"temperature": 0.1}), + ), + ( + json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]), + json!({}), + ), + ] { + assert_eq!( + decline_reason( + "anthropic/claude-sonnet-4-5", + None, + messages.clone(), + params.clone() + ), + None, + "gate declined {messages}" + ); + prepare_chat_completions_call(request( + "anthropic/claude-sonnet-4-5", + None, + messages.clone(), + params, + )) + .unwrap_or_else(|error| panic!("prepare declined {messages}: {error}")); + } +} + +mod round_trip { + use super::*; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::{TcpListener, TcpStream}; + + use crate::chat_completions::chat_completions; + + async fn read_http_request(socket: &mut TcpStream) -> String { + let mut request = Vec::new(); + let mut buffer = [0_u8; 1024]; + let header_end = loop { + let n = socket.read(&mut buffer).await.expect("reads request"); + if n == 0 { + break request.len(); + } + request.extend_from_slice(&buffer[..n]); + if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") { + break position + 4; + } + }; + let headers = String::from_utf8_lossy(&request[..header_end]); + let content_length = headers + .lines() + .find_map(|line| { + let (name, value) = line.split_once(':')?; + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().ok()) + .flatten() + }) + .unwrap_or(0); + while request.len().saturating_sub(header_end) < content_length { + let n = socket.read(&mut buffer).await.expect("reads body"); + if n == 0 { + break; + } + request.extend_from_slice(&buffer[..n]); + } + String::from_utf8(request).expect("request is utf8") + } + + fn http_response(status: &str, body: &str) -> String { + format!( + "HTTP/1.1 {status}\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", + body.len(), + body + ) + } + + /// Serve one request from a stub upstream and hand back what it received. + async fn serve_once( + status: &'static str, + body: &'static str, + ) -> (String, tokio::task::JoinHandle) { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + let port = listener.local_addr().expect("addr").port(); + let handle = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts"); + let received = read_http_request(&mut socket).await; + socket + .write_all(http_response(status, body).as_bytes()) + .await + .expect("writes response"); + socket.flush().await.expect("flushes"); + received + }); + (format!("http://127.0.0.1:{port}/v1/messages"), handle) + } + + fn call(api_base: &str, messages: Value, params: Value) -> ChatCompletionsRequest<'_> { + ChatCompletionsRequest { + model: "anthropic/claude-sonnet-4-5", + messages, + optional_params: match params { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + }, + api_key: Some("sk-test"), + api_base: Some(api_base), + custom_llm_provider: None, + extra_headers: None, + timeout: Some(std::time::Duration::from_secs(10)), + } + } + + const GOOD_BODY: &str = r#"{"id":"msg_1","type":"message","role":"assistant","model":"claude-sonnet-4-5-20260101","content":[{"type":"text","text":"hello"}],"stop_reason":"end_turn","stop_sequence":null,"usage":{"input_tokens":11,"output_tokens":4}}"#; + + #[tokio::test] + async fn round_trip_sends_the_translated_body_and_normalizes_the_response() { + let (api_base, handle) = serve_once("200 OK", GOOD_BODY).await; + let response = chat_completions(call( + &api_base, + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"} + ]), + json!({"max_tokens": 16}), + )) + .await + .expect("call succeeds"); + + let received = handle.await.expect("server task"); + let sent: Value = serde_json::from_str( + received + .split_once("\r\n\r\n") + .expect("request has a body") + .1, + ) + .expect("body is json"); + assert_eq!( + sent["messages"], + json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]) + ); + assert_eq!( + sent["system"], + json!([{"type": "text", "text": "be terse"}]) + ); + assert_eq!(sent["max_tokens"], json!(16)); + assert!(received.to_lowercase().contains("x-api-key: sk-test")); + + assert_eq!( + response.choices[0].message.content.as_deref(), + Some("hello") + ); + assert_eq!(response.usage.total_tokens, 15); + } + + #[tokio::test] + async fn a_response_it_cannot_normalize_is_reported_as_already_sent() { + // The provider was called and billed, so the host must not retry this + // on its own path. `MissingField` here would read as a pre-send + // decline and be retried; `InvalidResponse` cannot. + const NO_USAGE: &str = + r#"{"model":"m","content":[{"type":"text","text":"hi"}],"stop_reason":"end_turn"}"#; + let (api_base, handle) = serve_once("200 OK", NO_USAGE).await; + let err = chat_completions(call( + &api_base, + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + )) + .await + .expect_err("response cannot be normalized"); + handle.await.expect("server task"); + assert!( + matches!(err, CoreError::InvalidResponse(_)), + "expected a post-send error, got {err:?}" + ); + } + + #[tokio::test] + async fn a_tool_use_block_in_the_response_is_also_reported_as_already_sent() { + const TOOL_USE: &str = r#"{"model":"m","content":[{"type":"tool_use","id":"t","name":"f","input":{}}],"stop_reason":"tool_use","usage":{"input_tokens":1,"output_tokens":1}}"#; + let (api_base, handle) = serve_once("200 OK", TOOL_USE).await; + let err = chat_completions(call( + &api_base, + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + )) + .await + .expect_err("response cannot be normalized"); + handle.await.expect("server task"); + assert!( + matches!(err, CoreError::InvalidResponse(_)), + "expected a post-send error, got {err:?}" + ); + } + + #[tokio::test] + async fn an_upstream_error_status_keeps_its_code() { + let (api_base, handle) = + serve_once("429 Too Many Requests", r#"{"error":"slow down"}"#).await; + let err = chat_completions(call( + &api_base, + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + )) + .await + .expect_err("upstream rejects"); + handle.await.expect("server task"); + assert!( + matches!(err, CoreError::Http { status: 429, .. }), + "expected a 429, got {err:?}" + ); + } + + #[tokio::test] + async fn a_connection_that_is_never_established_declines_instead_of_failing() { + // Nothing was sent, so nothing was billed and the host can still serve + // the request. Classing this with the post-send failures would turn a + // recoverable fallback into a user-facing error on exactly the + // deployments whose transport is configured only on the Python client. + let port = { + let listener = TcpListener::bind("127.0.0.1:0").await.expect("binds"); + listener.local_addr().expect("has an address").port() + // Dropped here, so the port is closed and the connect is refused. + }; + let err = chat_completions(call( + &format!("http://127.0.0.1:{port}/v1/messages"), + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + )) + .await + .expect_err("nothing is listening"); + assert!( + matches!(err, CoreError::Connect(_)), + "expected a pre-send connect failure, got {err:?}" + ); + } + + #[test] + fn response_errors_collapse_to_one_variant_that_can_only_mean_already_sent() { + use crate::chat_completions::handler::as_response_error; + + for original in [ + CoreError::MissingField("usage"), + CoreError::Unsupported("non-text response content block"), + CoreError::InvalidRequest("whatever".to_string()), + CoreError::Auth("whatever".to_string()), + ] { + let label = format!("{original:?}"); + assert!( + matches!(as_response_error(original), CoreError::InvalidResponse(_)), + "{label} must not stay retryable once the provider has answered" + ); + } + // An upstream status is already unambiguous, so it survives intact. + assert!(matches!( + as_response_error(CoreError::Http { + status: 500, + body: "boom".to_string() + }), + CoreError::Http { status: 500, .. } + )); + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/transformation.rs b/litellm-rust/crates/core/src/chat_completions/transformation.rs new file mode 100644 index 00000000000..a30ce9dc77c --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/transformation.rs @@ -0,0 +1,155 @@ +use serde_json::{Map, Value}; + +use crate::error::CoreResult; + +use super::types::{ + ChatCompletionsResponse, ChatMessage, ChatMessageContent, ProviderChatRequestData, + ProviderChatResponseData, +}; + +/// How the upstream call is authenticated. API-key strategies are resolved in +/// `prepare`; SigV4 needs the serialized body, so the handler signs it. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum ChatCompletionsAuth { + Header { name: &'static str, value: String }, + Bearer { token: String }, + AwsSigV4 { region: String }, +} + +/// Why a request cannot be served by the Rust path. +/// +/// The core declines rather than guessing: the host turns this into a +/// transparent fallback to the Python implementation, which covers the full +/// surface. Acceptance is an allowlist, so a parameter or message shape the +/// core has never seen declines by construction instead of being translated +/// wrong. +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct Unsupported(pub &'static str); + +pub const STREAM_PARAM: &str = "stream"; + +/// Message fields that carry no meaning for the upstream body, so their +/// presence does not make a request untranslatable. +const IGNORABLE_MESSAGE_FIELDS: &[&str] = &["name"]; + +pub trait ChatCompletionsProviderConfig: Sync { + fn complete_url( + &self, + api_base: Option<&str>, + model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult; + + fn auth( + &self, + api_key: Option<&str>, + model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult; + + fn default_headers(&self) -> &'static [(&'static str, &'static str)] { + &[("content-type", "application/json")] + } + + /// Whether an auth header the caller already supplied is the credential this + /// request should authenticate with, so the resolved one is not applied. + /// + /// Defaults to false: the deployment's credential outranks anything + /// forwarded, which is what every provider wants for its own auth header. + /// A provider overrides this only for a scheme it hands off to entirely. + fn defers_to_forwarded_auth(&self, _headers: &[(String, String)]) -> bool { + false + } + + /// Provider parameter names (post-mapping) the Rust path knows how to place + /// in the upstream body. Anything outside this set declines the request. + fn supported_params(&self) -> &'static [&'static str]; + + /// Parameters consumed as call configuration (credentials, endpoints) + /// rather than placed in the body. Accepted, never serialized. + fn config_params(&self) -> &'static [&'static str] { + &[] + } + + fn unsupported_reason( + &self, + messages: &[ChatMessage], + optional_params: &Map, + ) -> Option { + unsupported_param( + self.supported_params(), + self.config_params(), + optional_params, + ) + .or_else(|| messages.iter().find_map(unsupported_message)) + } + + fn transform_request( + &self, + model: &str, + messages: Vec, + optional_params: Map, + ) -> CoreResult; + + fn transform_response( + &self, + model: &str, + response: ProviderChatResponseData, + ) -> CoreResult; +} + +pub fn unsupported_param( + supported: &'static [&'static str], + config: &'static [&'static str], + optional_params: &Map, +) -> Option { + if optional_params + .get(STREAM_PARAM) + .and_then(Value::as_bool) + .unwrap_or(false) + { + return Some(Unsupported("streaming")); + } + optional_params + .keys() + .any(|key| { + key != STREAM_PARAM + && !supported.contains(&key.as_str()) + && !config.contains(&key.as_str()) + }) + .then_some(Unsupported("unrecognized request parameter")) +} + +/// Message shapes the core can translate faithfully: text content, either a +/// plain string or a non-empty list of parts that are all +/// `{"type": "text", "text": ...}`. Tool calls, tool results, and multimodal +/// parts decline so Python's fuller translation handles them. +pub fn unsupported_message(message: &ChatMessage) -> Option { + if message + .extra + .keys() + .any(|key| !IGNORABLE_MESSAGE_FIELDS.contains(&key.as_str())) + { + return Some(Unsupported("unrecognized message field")); + } + if !matches!(message.role.as_str(), "system" | "user" | "assistant") { + return Some(Unsupported("unrecognized message role")); + } + match &message.content { + None => Some(Unsupported("message without content")), + Some(ChatMessageContent::Text(_)) => None, + Some(ChatMessageContent::Parts(parts)) if parts.is_empty() => { + Some(Unsupported("message without content")) + } + Some(ChatMessageContent::Parts(parts)) => parts + .iter() + .any(|part| { + part.get("type").and_then(Value::as_str) != Some("text") + || part.get("text").and_then(Value::as_str).is_none() + || part.as_object().is_some_and(|object| object.len() != 2) + }) + .then_some(Unsupported("non-text message content")), + } +} diff --git a/litellm-rust/crates/core/src/chat_completions/types.rs b/litellm-rust/crates/core/src/chat_completions/types.rs new file mode 100644 index 00000000000..35dd543a986 --- /dev/null +++ b/litellm-rust/crates/core/src/chat_completions/types.rs @@ -0,0 +1,112 @@ +use std::time::Duration; + +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; + +use super::transformation::{ChatCompletionsAuth, ChatCompletionsProviderConfig}; + +/// A `/chat/completions` call as it crosses into the core. +/// +/// `optional_params` arrives already mapped to the provider's own parameter +/// names by the host, exactly as the messages route receives an already +/// Anthropic-shaped body. The core owns the conversation translation, the +/// provider call, and the response normalization. +pub struct ChatCompletionsRequest<'a> { + pub model: &'a str, + pub messages: Value, + pub optional_params: Map, + pub api_key: Option<&'a str>, + pub api_base: Option<&'a str>, + pub custom_llm_provider: Option<&'a str>, + pub extra_headers: Option>, + pub timeout: Option, +} + +pub(super) struct ProviderChatCompletionsRequest { + pub(super) model: String, + pub(super) config: &'static dyn ChatCompletionsProviderConfig, + pub(super) url: String, + pub(super) body: Value, + pub(super) upstream_headers: Vec<(String, String)>, + pub(super) auth: ChatCompletionsAuth, + #[cfg_attr(not(feature = "bedrock-auth"), allow(dead_code))] + pub(super) optional_params: Map, + pub(super) timeout: Option, +} + +/// The provider-shaped request body a config produces. Named rather than a bare +/// `Value` so the transform contract stays a typed one, mirroring +/// [`crate::audio_transcription::types::AudioTranscriptionRequestData`]. +pub struct ProviderChatRequestData { + pub body: Value, +} + +/// The raw provider response body handed back to a config for normalization. +pub struct ProviderChatResponseData { + pub body: Value, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum ChatMessageContent { + Text(String), + Parts(Vec), +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChatMessage { + pub role: String, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(default, skip_serializing_if = "Option::is_none")] + pub name: Option, + #[serde(flatten)] + pub extra: Map, +} + +/// OpenAI `usage`, including the `prompt_tokens_details` split LiteLLM's Python +/// path reports so cost tracking sees the same numbers on either path. +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct PromptTokensDetails { + pub cached_tokens: u64, + pub cache_creation_tokens: u64, + pub text_tokens: u64, +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub struct ChatCompletionsUsage { + pub prompt_tokens: u64, + pub completion_tokens: u64, + pub total_tokens: u64, + pub prompt_tokens_details: PromptTokensDetails, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChatCompletionsChoiceMessage { + pub role: String, + // Whether an empty turn is `None` or `""` is the provider's choice, not a + // shared invariant: Anthropic's transform ends on `merged_text or None` + // while Converse assigns the joined string unconditionally. Each config + // mirrors its own, so keep this optional and serialize it even when None. + pub content: Option, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChatCompletionsChoice { + pub index: u64, + pub message: ChatCompletionsChoiceMessage, + pub finish_reason: String, +} + +/// The normalized response handed back to the host. +/// +/// There is deliberately no `id`: Python mints the `chatcmpl-…` id on the +/// `ModelResponse` it already created, and echoing the provider's own id here +/// would change it. Pinned by `response_carries_no_id` in `tests.rs`. +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct ChatCompletionsResponse { + pub created: u64, + pub model: String, + pub choices: Vec, + pub usage: ChatCompletionsUsage, +} diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index caada1d98b0..e1ac0a4fc8f 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -12,8 +12,30 @@ pub(crate) const MESSAGES_CONNECT_TIMEOUT_SECS: u64 = 10; /// Max characters of an upstream error body echoed across the call boundary /// before truncation, so provider bodies are bounded and data-minimized. -pub(crate) const MESSAGES_ERROR_BODY_MAX_CHARS: usize = 256; +pub(crate) const UPSTREAM_ERROR_BODY_MAX_CHARS: usize = 256; /// Provider name used for Anthropic Messages when a deployment's provider model /// does not carry an explicit provider prefix. pub const ANTHROPIC_MESSAGES_PROVIDER: &str = "anthropic"; + +/// Prefix identifying an Anthropic OAuth token. Mirrors Python's +/// `ANTHROPIC_OAUTH_TOKEN_PREFIX`, which is what makes `validate_environment` +/// authenticate with `authorization` and drop `x-api-key` entirely. +pub(crate) const ANTHROPIC_OAUTH_TOKEN_PREFIX: &str = "sk-ant-oat"; + +/// Full-request timeout ceiling for chat completions provider calls, in +/// seconds. Mirrors the Python chat completions default. +pub(crate) const CHAT_COMPLETIONS_TIMEOUT_SECS: u64 = 600; + +/// Connect timeout for chat completions provider calls, in seconds. +pub(crate) const CHAT_COMPLETIONS_CONNECT_TIMEOUT_SECS: u64 = 10; + +/// `object` field every non-streaming chat completion response carries. +pub const CHAT_COMPLETION_OBJECT: &str = "chat.completion"; + +/// Placeholder Python substitutes for empty or whitespace-only message text, +/// which Anthropic and Bedrock both reject. Must match +/// `_EMPTY_TEXT_PLACEHOLDER` in +/// `litellm/litellm_core_utils/prompt_templates/factory.py`. +pub const EMPTY_TEXT_PLACEHOLDER: &str = + "[System: Empty message content sanitised to satisfy protocol]"; diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index c2b08eee0c0..739532f8cb5 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -23,8 +23,19 @@ pub enum CoreError { Http { status: u16, body: String }, #[error("upstream network error: {0}")] Network(String), + /// The provider was never reached: DNS, TCP, TLS or proxy setup failed + /// before any byte of the request went out. Nothing was billed, so a host + /// that keeps a reference implementation can serve the request itself. + /// A timeout is deliberately not this, since the provider may have received + /// and answered the request already. + #[error("could not reach the provider: {0}")] + Connect(String), #[error("routing error: {0}")] Routing(String), + /// The request is outside the surface this route covers in Rust. Hosts that + /// keep a reference implementation treat this as "fall back", not "fail". + #[error("unsupported by the rust path: {0}")] + Unsupported(&'static str), } pub fn json_type_name(value: &serde_json::Value) -> &'static str { diff --git a/litellm-rust/crates/core/src/http_utils.rs b/litellm-rust/crates/core/src/http_utils.rs new file mode 100644 index 00000000000..c541f50275b --- /dev/null +++ b/litellm-rust/crates/core/src/http_utils.rs @@ -0,0 +1,112 @@ +//! Header and upstream-body helpers shared by every route module. + +use serde_json::{Map, Value}; + +use crate::constants::UPSTREAM_ERROR_BODY_MAX_CHARS; +use crate::error::{CoreError, CoreResult, json_type_name}; + +/// Bound an upstream error body before it crosses a host boundary, so provider +/// bodies stay data-minimized. +pub fn truncate_error_body(body: &str) -> String { + if body.chars().count() <= UPSTREAM_ERROR_BODY_MAX_CHARS { + return body.to_string(); + } + let truncated: String = body.chars().take(UPSTREAM_ERROR_BODY_MAX_CHARS).collect(); + format!("{truncated}... (truncated)") +} + +pub fn string_headers( + context: &'static str, + extra_headers: Option>, +) -> CoreResult> { + extra_headers + .unwrap_or_default() + .into_iter() + .map(|(key, value)| { + value + .as_str() + .map(|value| (key.clone(), value.to_string())) + .ok_or_else(|| { + CoreError::InvalidRequest(format!( + "{context} extra_headers.{key} must be a string, got {}", + json_type_name(&value) + )) + }) + }) + .collect() +} + +pub fn has_header(headers: &[(String, String)], name: &str) -> bool { + headers + .iter() + .any(|(key, _)| key.eq_ignore_ascii_case(name)) +} + +pub fn has_bearer_auth(headers: &[(String, String)]) -> bool { + headers.iter().any(|(name, value)| { + if !name.eq_ignore_ascii_case("authorization") { + return false; + } + let value = value.trim(); + value.len() > 7 + && value[..7].eq_ignore_ascii_case("bearer ") + && !value[7..].trim().is_empty() + }) +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::json; + + #[test] + fn truncate_leaves_short_bodies_untouched() { + assert_eq!(truncate_error_body("short"), "short"); + } + + #[test] + fn truncate_bounds_long_bodies_by_characters() { + let body = "\u{00e9}".repeat(UPSTREAM_ERROR_BODY_MAX_CHARS + 10); + let truncated = truncate_error_body(&body); + assert!(truncated.ends_with("... (truncated)")); + assert_eq!( + truncated.chars().count(), + UPSTREAM_ERROR_BODY_MAX_CHARS + "... (truncated)".chars().count() + ); + } + + #[test] + fn string_headers_rejects_non_string_values() { + let headers = Map::from_iter([("x-trace".to_string(), json!(7))]); + let err = string_headers("chat completions", Some(headers)).expect_err("non-string value"); + assert_eq!( + err, + CoreError::InvalidRequest( + "chat completions extra_headers.x-trace must be a string, got number".to_string() + ) + ); + } + + #[test] + fn header_lookup_is_case_insensitive() { + let headers = vec![("X-Api-Key".to_string(), "k".to_string())]; + assert!(has_header(&headers, "x-api-key")); + assert!(!has_header(&headers, "authorization")); + } + + #[test] + fn bearer_detection_requires_a_non_empty_token() { + assert!(has_bearer_auth(&[( + "Authorization".to_string(), + "Bearer abc".to_string() + )])); + assert!(!has_bearer_auth(&[( + "Authorization".to_string(), + "Bearer ".to_string() + )])); + assert!(!has_bearer_auth(&[( + "Authorization".to_string(), + "Basic abc".to_string() + )])); + } +} diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 51ea19750ea..dce4a425ea0 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,8 +1,10 @@ pub mod audio_transcription; pub mod caching; pub mod call_lifecycle; +pub mod chat_completions; pub mod constants; pub mod error; +pub mod http_utils; pub mod messages; pub mod ocr; pub mod providers; diff --git a/litellm-rust/crates/core/src/messages/common_utils.rs b/litellm-rust/crates/core/src/messages/common_utils.rs index 9dcfcaa71e3..a14dffbc1fe 100644 --- a/litellm-rust/crates/core/src/messages/common_utils.rs +++ b/litellm-rust/crates/core/src/messages/common_utils.rs @@ -1,19 +1,15 @@ use serde_json::{Map, Value}; -use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS; -use crate::error::{CoreError, CoreResult, json_type_name}; +use crate::error::CoreResult; +use crate::http_utils::string_headers as shared_string_headers; use crate::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; use crate::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; use super::transformation::AnthropicMessagesProviderConfig; -pub(super) fn truncate_error_body(body: &str) -> String { - if body.chars().count() <= MESSAGES_ERROR_BODY_MAX_CHARS { - return body.to_string(); - } - let truncated: String = body.chars().take(MESSAGES_ERROR_BODY_MAX_CHARS).collect(); - format!("{truncated}... (truncated)") -} +pub(super) use crate::http_utils::{has_bearer_auth, has_header, truncate_error_body}; + +const HEADER_CONTEXT: &str = "messages"; pub(super) fn messages_provider_config( provider: &str, @@ -28,37 +24,5 @@ pub(super) fn messages_provider_config( pub(super) fn string_headers( extra_headers: Option>, ) -> CoreResult> { - extra_headers - .unwrap_or_default() - .into_iter() - .map(|(key, value)| { - value - .as_str() - .map(|value| (key.clone(), value.to_string())) - .ok_or_else(|| { - CoreError::InvalidRequest(format!( - "messages extra_headers.{key} must be a string, got {}", - json_type_name(&value) - )) - }) - }) - .collect() -} - -pub(super) fn has_header(headers: &[(String, String)], name: &str) -> bool { - headers - .iter() - .any(|(key, _)| key.eq_ignore_ascii_case(name)) -} - -pub(super) fn has_bearer_auth(headers: &[(String, String)]) -> bool { - headers.iter().any(|(name, value)| { - if !name.eq_ignore_ascii_case("authorization") { - return false; - } - let value = value.trim(); - value.len() > 7 - && value[..7].eq_ignore_ascii_case("bearer ") - && !value[7..].trim().is_empty() - }) + shared_string_headers(HEADER_CONTEXT, extra_headers) } diff --git a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/mod.rs b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/mod.rs new file mode 100644 index 00000000000..f239b6921fa --- /dev/null +++ b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/mod.rs @@ -0,0 +1 @@ +pub mod transformation; diff --git a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs new file mode 100644 index 00000000000..4534ac0182c --- /dev/null +++ b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/tests.rs @@ -0,0 +1,444 @@ +use super::*; +use serde_json::json; + +fn messages(value: Value) -> Vec { + serde_json::from_value(value).expect("valid messages") +} + +fn params(value: Value) -> Map { + match value { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + } +} + +fn transform(model: &str, msgs: Value, opts: Value) -> Value { + ANTHROPIC_CHAT_COMPLETIONS_CONFIG + .transform_request(model, messages(msgs), params(opts)) + .expect("request transforms") + .body +} + +fn transform_response(body: Value) -> CoreResult { + ANTHROPIC_CHAT_COMPLETIONS_CONFIG + .transform_response("claude-sonnet-4-5", ProviderChatResponseData { body }) +} + +fn reason(msgs: Value, opts: Value) -> Option { + ANTHROPIC_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), ¶ms(opts)) +} + +#[test] +fn builds_the_messages_body_python_builds() { + let body = transform( + "claude-sonnet-4-5", + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"} + ]), + json!({"max_tokens": 128, "temperature": 0.2}), + ); + assert_eq!( + body, + json!({ + "model": "claude-sonnet-4-5", + "messages": [ + {"role": "user", "content": [{"type": "text", "text": "hi"}]} + ], + "system": [{"type": "text", "text": "be terse"}], + "max_tokens": 128, + "temperature": 0.2 + }) + ); +} + +#[test] +fn omits_system_when_no_system_message_is_present() { + let body = transform( + "claude-sonnet-4-5", + json!([{"role": "user", "content": "hi"}]), + json!({"max_tokens": 16}), + ); + assert!(body.get("system").is_none()); +} + +#[test] +fn merges_consecutive_turns_and_wraps_every_text_in_a_block() { + let body = transform( + "claude-sonnet-4-5", + json!([ + {"role": "user", "content": "one"}, + {"role": "user", "content": [{"type": "text", "text": "two"}]}, + {"role": "assistant", "content": "ack"} + ]), + json!({"max_tokens": 16}), + ); + assert_eq!( + body["messages"], + json!([ + {"role": "user", "content": [ + {"type": "text", "text": "one"}, + {"type": "text", "text": "two"} + ]}, + {"role": "assistant", "content": [{"type": "text", "text": "ack"}]} + ]) + ); +} + +#[test] +fn right_strips_a_trailing_assistant_prefill_like_python() { + let body = transform( + "claude-sonnet-4-5", + json!([ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "Argentina "} + ]), + json!({"max_tokens": 16}), + ); + assert_eq!( + body["messages"][1]["content"][0]["text"], + json!("Argentina") + ); +} + +#[test] +fn passes_every_supported_param_through_untouched() { + let body = transform( + "claude-sonnet-4-5", + json!([{"role": "user", "content": "hi"}]), + json!({ + "max_tokens": 64, + "temperature": 0.1, + "top_p": 0.9, + "stop_sequences": ["STOP"] + }), + ); + assert_eq!(body["max_tokens"], json!(64)); + assert_eq!(body["temperature"], json!(0.1)); + assert_eq!(body["top_p"], json!(0.9)); + assert_eq!(body["stop_sequences"], json!(["STOP"])); +} + +#[test] +fn declines_top_k_because_python_gates_it_by_model_below_this_point() { + // `temperature` and `top_p` arrive already resolved, because + // `map_openai_params` applies `_apply_sampling_param` to them before the + // gate runs. `top_k` bypasses that and is gated inside `transform_request`, + // the function this route replaces, so forwarding it would send `top_k` to + // a model that removed sampling params and take a 400 after the call, where + // Python drops it and succeeds. + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({"top_k": 40}) + ), + Some(Unsupported("unrecognized request parameter")) + ); +} + +#[test] +fn declines_streaming_before_anything_else() { + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({"stream": true, "max_tokens": 16}) + ), + Some(Unsupported("streaming")) + ); +} + +#[test] +fn accepts_an_explicit_stream_false() { + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({"stream": false, "max_tokens": 16}) + ), + None + ); +} + +#[test] +fn declines_any_param_outside_the_allowlist() { + for param in [ + json!({"tools": []}), + json!({"tool_choice": {"type": "auto"}}), + json!({"thinking": {"type": "enabled"}}), + json!({"system": "injected"}), + json!({"metadata": {"user_id": "u1"}}), + json!({"output_config": {"effort": "high"}}), + ] { + assert_eq!( + reason(json!([{"role": "user", "content": "hi"}]), param.clone()), + Some(Unsupported("unrecognized request parameter")), + "expected {param} to decline" + ); + } +} + +#[test] +fn declines_tool_calls_tool_results_and_multimodal_content() { + assert_eq!( + reason( + json!([ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": null, "tool_calls": [ + {"id": "c1", "type": "function", + "function": {"name": "f", "arguments": "{}"}} + ]} + ]), + json!({}) + ), + Some(Unsupported("unrecognized message field")) + ); + assert_eq!( + reason( + json!([ + {"role": "user", "content": "hi"}, + {"role": "tool", "tool_call_id": "c1", "content": "ok"} + ]), + json!({}) + ), + Some(Unsupported("unrecognized message field")) + ); + assert_eq!( + reason( + json!([{"role": "user", "content": [ + {"type": "image_url", "image_url": {"url": "https://x/y.png"}} + ]}]), + json!({}) + ), + Some(Unsupported("non-text message content")) + ); + assert_eq!( + reason( + json!([{"role": "user", "content": [ + {"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}} + ]}]), + json!({}) + ), + Some(Unsupported("non-text message content")) + ); +} + +#[test] +fn declines_a_message_whose_content_list_is_empty() { + // An empty list passes every per-part check, so without this it would reach + // the provider as an empty `content` array and fail after the call rather + // than declining to Python before it. + assert_eq!( + reason(json!([{"role": "user", "content": []}]), json!({})), + Some(Unsupported("message without content")) + ); + assert_eq!( + reason( + json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]), + json!({}) + ), + None + ); +} + +#[test] +fn declines_a_conversation_that_does_not_open_on_a_user_turn() { + assert_eq!( + reason( + json!([ + {"role": "system", "content": "be terse"}, + {"role": "assistant", "content": "prefill"} + ]), + json!({}) + ), + Some(Unsupported("conversation does not open on a user turn")) + ); +} + +#[test] +fn accepts_a_plain_text_conversation() { + assert_eq!( + reason( + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "hello"}, + {"role": "user", "content": [{"type": "text", "text": "again"}]} + ]), + json!({"max_tokens": 16, "temperature": 0.5}) + ), + None + ); +} + +#[test] +fn normalizes_a_text_response_into_openai_shape() { + let response = transform_response(json!({ + "id": "msg_123", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20260101", + "content": [{"type": "text", "text": "hello"}, {"type": "text", "text": " there"}], + "stop_reason": "end_turn", + "stop_sequence": null, + "usage": {"input_tokens": 11, "output_tokens": 4} + })) + .expect("response transforms"); + + assert_eq!(response.model, "claude-sonnet-4-5-20260101"); + assert_eq!(response.choices.len(), 1); + assert_eq!(response.choices[0].index, 0); + assert_eq!(response.choices[0].message.role, "assistant"); + assert_eq!( + response.choices[0].message.content.as_deref(), + Some("hello there") + ); + assert_eq!(response.choices[0].finish_reason, "stop"); + assert_eq!(response.usage.prompt_tokens, 11); + assert_eq!(response.usage.completion_tokens, 4); + assert_eq!(response.usage.total_tokens, 15); +} + +#[test] +fn folds_cache_tokens_into_prompt_tokens_like_python() { + let response = transform_response(json!({ + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "hi"}], + "stop_reason": "end_turn", + "usage": { + "input_tokens": 10, + "output_tokens": 2, + "cache_read_input_tokens": 5, + "cache_creation_input_tokens": 3 + } + })) + .expect("response transforms"); + assert_eq!(response.usage.prompt_tokens, 18); + assert_eq!(response.usage.total_tokens, 20); + assert_eq!(response.usage.prompt_tokens_details.cached_tokens, 5); + assert_eq!( + response.usage.prompt_tokens_details.cache_creation_tokens, + 3 + ); + assert_eq!(response.usage.prompt_tokens_details.text_tokens, 10); +} + +#[test] +fn maps_max_tokens_stop_reason_to_length() { + let response = transform_response(json!({ + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "hi"}], + "stop_reason": "max_tokens", + "usage": {"input_tokens": 1, "output_tokens": 1} + })) + .expect("response transforms"); + assert_eq!(response.choices[0].finish_reason, "length"); +} + +#[test] +fn a_refusal_returns_the_completion_python_returns() { + // `refusal` is a stop_reason, not a content block type, so the content is + // ordinary text and this normalizes rather than declining. Python maps it + // to content_filter in _FINISH_REASON_MAP and returns the completion. + let response = transform_response(json!({ + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "I can't help with that."}], + "stop_reason": "refusal", + "usage": {"input_tokens": 9, "output_tokens": 6} + })) + .expect("a refusal still transforms"); + assert_eq!(response.choices[0].finish_reason, "content_filter"); + assert_eq!( + response.choices[0].message.content.as_deref(), + Some("I can't help with that.") + ); +} + +#[test] +fn reports_no_content_rather_than_an_empty_string() { + let response = transform_response(json!({ + "model": "claude-sonnet-4-5", + "content": [], + "stop_reason": "end_turn", + "usage": {"input_tokens": 1, "output_tokens": 0} + })) + .expect("response transforms"); + assert_eq!(response.choices[0].message.content, None); +} + +#[test] +fn response_carries_no_id_so_python_keeps_its_chatcmpl_id() { + let response = transform_response(json!({ + "id": "msg_should_not_leak", + "model": "claude-sonnet-4-5", + "content": [{"type": "text", "text": "hi"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 1, "output_tokens": 1} + })) + .expect("response transforms"); + let value = serde_json::to_value(response).expect("serializable"); + assert!( + value.get("id").is_none(), + "the rust response must not carry an id, got {value}" + ); +} + +#[test] +fn declines_a_response_carrying_a_non_text_block() { + let err = transform_response(json!({ + "model": "claude-sonnet-4-5", + "content": [{"type": "tool_use", "id": "t1", "name": "f", "input": {}}], + "stop_reason": "tool_use", + "usage": {"input_tokens": 1, "output_tokens": 1} + })) + .expect_err("non-text block"); + assert_eq!( + err, + CoreError::Unsupported("non-text response content block") + ); +} + +#[test] +fn errors_on_a_response_missing_required_fields() { + assert_eq!( + transform_response(json!("nope")).expect_err("not an object"), + CoreError::InvalidResponse("messages response is not an object".to_string()) + ); + assert_eq!( + transform_response(json!({"model": "m", "usage": {}})).expect_err("no content"), + CoreError::MissingField("content") + ); + assert_eq!( + transform_response(json!({"model": "m", "content": []})).expect_err("no usage"), + CoreError::MissingField("usage") + ); + assert_eq!( + transform_response(json!({"content": [], "usage": {}})).expect_err("no model"), + CoreError::MissingField("model") + ); +} + +#[test] +fn resolves_the_messages_url_and_x_api_key_auth() { + let config = &ANTHROPIC_CHAT_COMPLETIONS_CONFIG; + assert_eq!( + config + .complete_url(None, "claude-sonnet-4-5", &Map::new(), &|_| None) + .expect("url builds"), + "https://api.anthropic.com/v1/messages" + ); + assert_eq!( + config + .auth(Some("sk-x"), "claude-sonnet-4-5", &Map::new(), &|_| None) + .expect("auth resolves"), + ChatCompletionsAuth::Header { + name: "x-api-key", + value: "sk-x".to_string() + } + ); + assert_eq!( + config.default_headers(), + &[ + ("anthropic-version", "2023-06-01"), + ("content-type", "application/json"), + ] + ); +} diff --git a/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs new file mode 100644 index 00000000000..3658642b539 --- /dev/null +++ b/litellm-rust/crates/core/src/providers/anthropic/chat_completions/transformation.rs @@ -0,0 +1,211 @@ +use serde_json::{Map, Value, json}; + +use crate::chat_completions::conversation::{Conversation, build_conversation}; +use crate::chat_completions::transformation::{ + ChatCompletionsAuth, ChatCompletionsProviderConfig, Unsupported, unsupported_message, + unsupported_param, +}; +use crate::chat_completions::types::{ + ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse, ChatMessage, + ProviderChatRequestData, ProviderChatResponseData, +}; +use crate::constants::ANTHROPIC_OAUTH_TOKEN_PREFIX; +use crate::error::{CoreError, CoreResult}; +use crate::providers::anthropic::messages::transformation::{ + complete_anthropic_url, resolve_anthropic_api_key, +}; + +use crate::chat_completions::response_utils::{finish_reason_for, unix_now, usage_from_parts}; + +/// Anthropic parameter names, post `map_openai_params`, that the Rust path can +/// place verbatim in the Messages body. +/// +/// `top_k` is deliberately absent even though the Messages API takes it. +/// `temperature` and `top_p` reach this gate already resolved, because +/// `map_openai_params` runs first and applies `_apply_sampling_param` to them. +/// `top_k` bypasses `map_openai_params` entirely, so Python applies that same +/// per-model gate inside `transform_request`, the function this route replaces. +/// Forwarding it would send `top_k` to a model that removed sampling params and +/// take a 400 after the call, where Python drops it and succeeds. +const SUPPORTED_PARAMS: &[&str] = &["max_tokens", "temperature", "top_p", "stop_sequences"]; + +pub struct AnthropicChatCompletionsConfig; + +pub const ANTHROPIC_CHAT_COMPLETIONS_CONFIG: AnthropicChatCompletionsConfig = + AnthropicChatCompletionsConfig; + +fn text_block(text: &str) -> Value { + json!({"type": "text", "text": text}) +} + +fn anthropic_body(model: &str, conversation: &Conversation, params: Map) -> Value { + let messages: Vec = conversation + .turns + .iter() + .map(|turn| { + json!({ + "role": turn.role.as_str(), + "content": turn.texts.iter().map(|text| text_block(text)).collect::>(), + }) + }) + .collect(); + + let system: Vec = conversation.system.iter().map(|s| text_block(s)).collect(); + + let body = Map::from_iter( + [ + ("model".to_string(), json!(model)), + ("messages".to_string(), json!(messages)), + ] + .into_iter() + // Python builds `{"model", "messages", **optional_params}` with + // `system` already folded into optional_params, so a caller-supplied + // key of the same name wins here too. + .chain((!system.is_empty()).then(|| ("system".to_string(), json!(system)))) + .chain(params), + ); + Value::Object(body) +} + +impl ChatCompletionsProviderConfig for AnthropicChatCompletionsConfig { + fn complete_url( + &self, + api_base: Option<&str>, + _model: &str, + _optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + Ok(complete_anthropic_url(api_base, env_lookup)) + } + + fn auth( + &self, + api_key: Option<&str>, + _model: &str, + _optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + Ok(ChatCompletionsAuth::Header { + name: "x-api-key", + value: resolve_anthropic_api_key(api_key, env_lookup)?, + }) + } + + fn default_headers(&self) -> &'static [(&'static str, &'static str)] { + &[ + ("anthropic-version", "2023-06-01"), + ("content-type", "application/json"), + ] + } + + /// An OAuth bearer is the whole credential: Python's `validate_environment` + /// authenticates with it and drops `x-api-key` rather than resolving one, so + /// the resolved key must not be applied over the top. Any other forwarded + /// `authorization` is unrelated to this header and does not defer, which is + /// also what Python does: it sends the deployment's `x-api-key` alongside. + fn defers_to_forwarded_auth(&self, headers: &[(String, String)]) -> bool { + headers.iter().any(|(name, value)| { + name.eq_ignore_ascii_case("authorization") + && value + .strip_prefix("Bearer ") + .is_some_and(|token| token.starts_with(ANTHROPIC_OAUTH_TOKEN_PREFIX)) + }) + } + + fn supported_params(&self) -> &'static [&'static str] { + SUPPORTED_PARAMS + } + + fn unsupported_reason( + &self, + messages: &[ChatMessage], + optional_params: &Map, + ) -> Option { + unsupported_param(SUPPORTED_PARAMS, &[], optional_params) + .or_else(|| messages.iter().find_map(unsupported_message)) + // Anthropic rejects a request whose first turn is not a user turn. + // Python only repairs that under `litellm.modify_params`, which the + // core cannot observe, so decline instead of guessing. + .or_else(|| { + (!build_conversation(messages).opens_on_user_turn()) + .then_some(Unsupported("conversation does not open on a user turn")) + }) + } + + fn transform_request( + &self, + model: &str, + messages: Vec, + optional_params: Map, + ) -> CoreResult { + Ok(ProviderChatRequestData { + body: anthropic_body(model, &build_conversation(&messages), optional_params), + }) + } + + fn transform_response( + &self, + _model: &str, + response: ProviderChatResponseData, + ) -> CoreResult { + let body = response.body.as_object().ok_or_else(|| { + CoreError::InvalidResponse("messages response is not an object".into()) + })?; + + let content = body + .get("content") + .and_then(Value::as_array) + .ok_or(CoreError::MissingField("content"))?; + // The route declines tool and thinking requests, so a non-text block + // means the response carries something this path never asked for. + // Decline rather than silently dropping it; the host falls back. + if content + .iter() + .any(|block| block.get("type").and_then(Value::as_str) != Some("text")) + { + return Err(CoreError::Unsupported("non-text response content block")); + } + let text: String = content + .iter() + .filter_map(|block| block.get("text").and_then(Value::as_str)) + .collect(); + + let usage = body + .get("usage") + .and_then(Value::as_object) + .ok_or(CoreError::MissingField("usage"))?; + let field = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0); + + Ok(ChatCompletionsResponse { + created: unix_now(), + model: body + .get("model") + .and_then(Value::as_str) + .ok_or(CoreError::MissingField("model"))? + .to_string(), + choices: vec![ChatCompletionsChoice { + index: 0, + message: ChatCompletionsChoiceMessage { + role: "assistant".to_string(), + content: (!text.is_empty()).then_some(text), + }, + finish_reason: finish_reason_for( + body.get("stop_reason") + .and_then(Value::as_str) + .unwrap_or(""), + ) + .to_string(), + }], + usage: usage_from_parts( + field("input_tokens"), + field("output_tokens"), + field("cache_read_input_tokens"), + field("cache_creation_input_tokens"), + ), + }) + } +} + +#[cfg(test)] +#[path = "tests.rs"] +mod tests; diff --git a/litellm-rust/crates/core/src/providers/anthropic/mod.rs b/litellm-rust/crates/core/src/providers/anthropic/mod.rs index ba63992f3cb..0bb20991ff7 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/mod.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/mod.rs @@ -1 +1,2 @@ +pub mod chat_completions; pub mod messages; diff --git a/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs b/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs index 86eb589e2c0..5e885734182 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs @@ -8,11 +8,8 @@ use crate::audio_transcription::types::{ }; use crate::error::{CoreError, CoreResult, json_type_name}; -use super::aws_base::AwsAuthConfig; -use super::constants::{ - AWS_REGION, AWS_REGION_NAME, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE, - DEFAULT_BEDROCK_REGION, -}; +pub use super::aws_base::{aws_auth_config, bedrock_model_id_and_region, resolve_bedrock_region}; +use super::constants::{BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE}; const SUPPORTED_PARAMS: &[&str] = &["language", "prompt", "temperature", "response_format"]; @@ -21,64 +18,6 @@ pub static BEDROCK_AUDIO_TRANSCRIPTION_CONFIG: BedrockAudioTranscriptionConfig = pub struct BedrockAudioTranscriptionConfig; -pub fn bedrock_model_id_and_region(model: &str) -> (String, Option) { - let mut stripped = model; - for prefix in ["bedrock/converse/", "bedrock/", "converse/"] { - if let Some(value) = stripped.strip_prefix(prefix) { - stripped = value; - break; - } - } - let mut region = None; - if let Some((candidate, remainder)) = stripped.split_once('/') - && is_bedrock_region(candidate) - { - region = Some(candidate.to_string()); - stripped = remainder; - } - for prefix in ["nova-2/", "nova/"] { - if let Some(value) = stripped.strip_prefix(prefix) { - stripped = value; - break; - } - } - if region.is_none() { - region = stripped - .strip_prefix("arn:") - .and_then(|value| value.split(':').nth(3)) - .filter(|value| !value.is_empty()) - .map(str::to_string); - } - (stripped.to_string(), region) -} - -fn is_bedrock_region(value: &str) -> bool { - value.len() > 3 - && value.contains('-') - && value - .chars() - .all(|char| char.is_ascii_alphanumeric() || char == '-') -} - -pub fn resolve_bedrock_region( - model_region: Option<&str>, - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, -) -> String { - if let Some(region) = optional_params - .get("aws_region_name") - .and_then(Value::as_str) - { - return region.to_string(); - } - if let Some(region) = model_region { - return region.to_string(); - } - env_lookup(AWS_REGION_NAME) - .or_else(|| env_lookup(AWS_REGION)) - .unwrap_or_else(|| DEFAULT_BEDROCK_REGION.to_string()) -} - fn audio_fields(audio: Value) -> CoreResult<(String, String)> { let object = audio.as_object().ok_or_else(|| CoreError::InvalidType { expected: "object", @@ -203,32 +142,6 @@ impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig { } } -pub fn aws_auth_config( - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, -) -> AwsAuthConfig { - let value = |key: &str| { - optional_params - .get(key) - .and_then(Value::as_str) - .map(str::to_string) - }; - let env = |key: &str| env_lookup(key); - AwsAuthConfig { - access_key_id: value("aws_access_key_id").or_else(|| env("AWS_ACCESS_KEY_ID")), - secret_access_key: value("aws_secret_access_key").or_else(|| env("AWS_SECRET_ACCESS_KEY")), - session_token: value("aws_session_token").or_else(|| env("AWS_SESSION_TOKEN")), - region_name: value("aws_region_name").or_else(|| env(AWS_REGION_NAME)), - session_name: value("aws_session_name").or_else(|| env("AWS_SESSION_NAME")), - profile_name: value("aws_profile_name").or_else(|| env("AWS_PROFILE_NAME")), - role_name: value("aws_role_name").or_else(|| env("AWS_ROLE_NAME")), - web_identity_token: value("aws_web_identity_token") - .or_else(|| env("AWS_WEB_IDENTITY_TOKEN")), - sts_endpoint: value("aws_sts_endpoint").or_else(|| env("AWS_STS_ENDPOINT")), - external_id: value("aws_external_id").or_else(|| env("AWS_EXTERNAL_ID")), - } -} - #[cfg(test)] mod tests { use super::*; diff --git a/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs b/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs index dc036a3cf21..b11639aa09b 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs @@ -12,13 +12,15 @@ use aws_sigv4::http_request::{ }; use aws_sigv4::sign::v4; use aws_smithy_runtime_api::client::identity::Identity; +use serde_json::{Map, Value}; use sha2::{Digest, Sha256}; use super::constants::{ - AWS_ACCESS_KEY_ID, AWS_EXTERNAL_ID, AWS_PROFILE_NAME, AWS_REGION_NAME, AWS_ROLE_ARN, - AWS_ROLE_NAME, AWS_SECRET_ACCESS_KEY, AWS_SESSION_NAME, AWS_SESSION_TOKEN, AWS_STS_ENDPOINT, - AWS_WEB_IDENTITY_TOKEN, AWS_WEB_IDENTITY_TOKEN_FILE, BEDROCK_SERVICE, - DEFAULT_SESSION_NAME_PREFIX, + AWS_ACCESS_KEY_ID, AWS_EXTERNAL_ID, AWS_PROFILE_NAME, AWS_REGION, AWS_REGION_NAME, + AWS_ROLE_ARN, AWS_ROLE_NAME, AWS_SECRET_ACCESS_KEY, AWS_SESSION_NAME, AWS_SESSION_TOKEN, + AWS_SIGNED_HEADER_NAMES, AWS_STS_ENDPOINT, AWS_WEB_IDENTITY_TOKEN, AWS_WEB_IDENTITY_TOKEN_FILE, + BEDROCK_SERVICE, DEFAULT_BEDROCK_REGION, DEFAULT_SESSION_NAME_PREFIX, + SIGV4_COMPUTED_HEADER_NAMES, }; const STATIC_CREDENTIALS_TTL: Duration = Duration::from_secs(3600 - 60); @@ -401,6 +403,33 @@ fn default_session_name() -> String { format!("{DEFAULT_SESSION_NAME_PREFIX}-{seconds}") } +/// The subset of `headers` SigV4 should cover. +/// +/// Python signs only these and reattaches the rest afterwards, so a forwarded +/// client header cannot change the canonical request and invalidate the +/// signature. Signing everything instead makes the request 403 on a header the +/// caller supplied, on a deployment that works on the Python path. +pub fn aws_signature_headers(headers: &BTreeMap) -> BTreeMap { + headers + .iter() + .filter(|(name, _)| { + let name = name.to_ascii_lowercase(); + AWS_SIGNED_HEADER_NAMES.contains(&name.as_str()) + || name.starts_with("x-amz-") + || name.starts_with("x-amzn-") + }) + .map(|(name, value)| (name.clone(), value.clone())) + .collect() +} + +/// Whether the signer produces `name` itself. +/// +/// Python's reattach loop skips these, so a caller-supplied copy never reaches +/// the wire next to the computed one. +pub fn is_sigv4_computed_header(name: &str) -> bool { + SIGV4_COMPUTED_HEADER_NAMES.contains(&name.to_ascii_lowercase().as_str()) +} + pub fn sign_bedrock_post( url: &str, body: &[u8], @@ -441,6 +470,121 @@ pub fn sign_bedrock_post( .collect()) } +/// Model-id and region parsing shared by every Bedrock route. +pub fn bedrock_model_id_and_region(model: &str) -> (String, Option) { + let mut stripped = model; + for prefix in ["bedrock/converse/", "bedrock/", "converse/"] { + if let Some(value) = stripped.strip_prefix(prefix) { + stripped = value; + break; + } + } + let mut region = None; + if let Some((candidate, remainder)) = stripped.split_once('/') + && is_bedrock_region(candidate) + { + region = Some(candidate.to_string()); + stripped = remainder; + } + for prefix in ["nova-2/", "nova/"] { + if let Some(value) = stripped.strip_prefix(prefix) { + stripped = value; + break; + } + } + if region.is_none() { + // Python splits the whole ARN and takes field 3, the region. Stripping + // `arn:` first shifts every field down one, so the region is field 2 + // here; field 3 is the account id. + region = stripped + .strip_prefix("arn:") + .and_then(|value| value.split(':').nth(2)) + .filter(|value| !value.is_empty()) + .map(str::to_string); + } + (stripped.to_string(), region) +} + +fn is_bedrock_region(value: &str) -> bool { + value.len() > 3 + && value.contains('-') + && value + .chars() + .all(|char| char.is_ascii_alphanumeric() || char == '-') +} + +pub fn resolve_bedrock_region( + model_region: Option<&str>, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, +) -> String { + if let Some(region) = optional_params + .get("aws_region_name") + .and_then(Value::as_str) + { + return region.to_string(); + } + if let Some(region) = model_region { + return region.to_string(); + } + env_lookup(AWS_REGION_NAME) + .or_else(|| env_lookup(AWS_REGION)) + .unwrap_or_else(|| DEFAULT_BEDROCK_REGION.to_string()) +} + +pub fn aws_auth_config( + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, +) -> AwsAuthConfig { + let value = |key: &str| { + optional_params + .get(key) + .and_then(Value::as_str) + .map(str::to_string) + }; + let env = |key: &str| env_lookup(key); + AwsAuthConfig { + access_key_id: value("aws_access_key_id").or_else(|| env("AWS_ACCESS_KEY_ID")), + secret_access_key: value("aws_secret_access_key").or_else(|| env("AWS_SECRET_ACCESS_KEY")), + session_token: value("aws_session_token").or_else(|| env("AWS_SESSION_TOKEN")), + region_name: value("aws_region_name").or_else(|| env(AWS_REGION_NAME)), + session_name: value("aws_session_name").or_else(|| env("AWS_SESSION_NAME")), + profile_name: value("aws_profile_name").or_else(|| env("AWS_PROFILE_NAME")), + role_name: value("aws_role_name").or_else(|| env("AWS_ROLE_NAME")), + web_identity_token: value("aws_web_identity_token") + .or_else(|| env("AWS_WEB_IDENTITY_TOKEN")), + sts_endpoint: value("aws_sts_endpoint").or_else(|| env("AWS_STS_ENDPOINT")), + external_id: value("aws_external_id").or_else(|| env("AWS_EXTERNAL_ID")), + } +} + +/// Credentials a host resolved through its own chain and handed down verbatim. +/// +/// A host with its own resolution (LiteLLM's Python `BaseAWSLLM`, which reads +/// profiles, STS and boto sessions) passes the result here so the core signs +/// with exactly those. Without this the core would re-derive from ambient +/// state, where an unrelated `AWS_ROLE_NAME` or `AWS_PROFILE_NAME` in the +/// environment outranks explicit keys in [`classify_auth`] and the two sides +/// would sign as different principals. +pub fn host_supplied_credentials(optional_params: &Map) -> Option { + let value = |key: &str| { + optional_params + .get(key) + .and_then(Value::as_str) + .map(str::trim) + .filter(|value| !value.is_empty()) + }; + let access_key_id = value("aws_access_key_id")?; + let secret_access_key = value("aws_secret_access_key")?; + Some(Credentials::new( + access_key_id, + secret_access_key, + value("aws_session_token").map(str::to_string), + None, + "litellm-host-supplied", + )) +} + #[cfg(test)] mod tests { use super::*; @@ -458,6 +602,18 @@ mod tests { ) } + #[test] + fn reads_the_region_field_of_a_model_arn_not_the_account_id() { + // Python's `_get_aws_region_from_model_arn` splits the whole ARN and + // takes field 3. Stripping `arn:` first shifts every field down one, so + // the region is field 2 here. Taking field 3 after the strip returns + // the account id, which is not a region at all. + let (_, region) = bedrock_model_id_and_region( + "bedrock/arn:aws:bedrock:us-west-2:123456789012:foundation-model/anthropic.claude-v2", + ); + assert_eq!(region.as_deref(), Some("us-west-2")); + } + #[test] fn classification_preserves_python_precedence() { let config = AwsAuthConfig { @@ -610,6 +766,52 @@ mod tests { )); } + #[test] + fn a_forwarded_client_header_is_not_folded_into_the_signature() { + // Python signs only the AWS header set, so a header a caller forwarded + // cannot change the canonical request. Signing it instead makes the + // request 403 the moment anything on the wire rewrites or drops it. + let (url, body, mut headers) = parity_inputs(); + headers.insert("x-request-id".to_string(), "abc-123".to_string()); + headers.insert("Accept-Encoding".to_string(), "gzip".to_string()); + headers.insert("x-amzn-trace-id".to_string(), "Root=1-abc".to_string()); + let signable = aws_signature_headers(&headers); + + assert!(!signable.contains_key("x-request-id")); + assert!(!signable.contains_key("Accept-Encoding")); + // The AWS-prefixed one is genuinely part of the signature. + assert!(signable.contains_key("x-amzn-trace-id")); + assert!(signable.contains_key("Content-Type")); + + let credentials = Credentials::new( + "AKIDEXAMPLE", + "wJalrXUtnFEMI/K7MDENG+bPxRfiCYEXAMPLEKEY", + None, + None, + "test", + ); + let signed = sign_bedrock_post( + &url, + &body, + &signable, + "us-east-1", + &credentials, + SystemTime::UNIX_EPOCH, + ) + .expect("signs"); + let authorization = signed + .get("Authorization") + .expect("carries an authorization header"); + assert!( + !authorization.contains("x-request-id"), + "forwarded header reached SignedHeaders: {authorization}" + ); + assert!( + !authorization.contains("accept-encoding"), + "forwarded header reached SignedHeaders: {authorization}" + ); + } + #[test] fn signing_matches_botocore_golden_vector() { let (url, body, headers) = parity_inputs(); diff --git a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/mod.rs b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/mod.rs new file mode 100644 index 00000000000..f239b6921fa --- /dev/null +++ b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/mod.rs @@ -0,0 +1 @@ +pub mod transformation; diff --git a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs new file mode 100644 index 00000000000..4b75dcb8e9d --- /dev/null +++ b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/tests.rs @@ -0,0 +1,580 @@ +use super::*; +use serde_json::json; + +fn messages(value: Value) -> Vec { + serde_json::from_value(value).expect("valid messages") +} + +fn params(value: Value) -> Map { + match value { + Value::Object(map) => map, + other => panic!("params must be an object, got {other}"), + } +} + +fn transform(msgs: Value, opts: Value) -> Value { + BEDROCK_CHAT_COMPLETIONS_CONFIG + .transform_request( + "anthropic.claude-sonnet-4-5-v1:0", + messages(msgs), + params(opts), + ) + .expect("request transforms") + .body +} + +fn transform_response(body: Value) -> CoreResult { + BEDROCK_CHAT_COMPLETIONS_CONFIG.transform_response( + "anthropic.claude-sonnet-4-5-v1:0", + ProviderChatResponseData { body }, + ) +} + +fn reason(msgs: Value, opts: Value) -> Option { + BEDROCK_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), ¶ms(opts)) +} + +#[test] +fn builds_the_converse_body_python_builds() { + let body = transform( + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"} + ]), + json!({"maxTokens": 128, "temperature": 0.2}), + ); + assert_eq!( + body, + json!({ + "inferenceConfig": {"maxTokens": 128, "temperature": 0.2}, + "messages": [{"role": "user", "content": [{"text": "hi"}]}], + "system": [{"text": "be terse"}] + }) + ); +} + +#[test] +fn always_emits_inference_config_even_when_empty() { + let body = transform(json!([{"role": "user", "content": "hi"}]), json!({})); + assert_eq!(body["inferenceConfig"], json!({})); + assert!(body.get("system").is_none()); +} + +#[test] +fn places_only_inference_params_in_inference_config() { + let body = transform( + json!([{"role": "user", "content": "hi"}]), + json!({ + "maxTokens": 64, + "temperature": 0.1, + "topP": 0.9, + "stopSequences": ["STOP"] + }), + ); + assert_eq!( + body["inferenceConfig"], + json!({"maxTokens": 64, "temperature": 0.1, "topP": 0.9, "stopSequences": ["STOP"]}) + ); + assert!(body.get("additionalModelRequestFields").is_none()); +} + +#[test] +fn merges_consecutive_user_turns_into_one_message() { + let body = transform( + json!([ + {"role": "user", "content": "one"}, + {"role": "user", "content": [{"type": "text", "text": "two"}]}, + {"role": "assistant", "content": "ack"}, + {"role": "user", "content": "three"} + ]), + json!({}), + ); + assert_eq!( + body["messages"], + json!([ + {"role": "user", "content": [{"text": "one"}, {"text": "two"}]}, + {"role": "assistant", "content": [{"text": "ack"}]}, + {"role": "user", "content": [{"text": "three"}]} + ]) + ); +} + +#[test] +fn declines_streaming() { + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({"stream": true}) + ), + Some(Unsupported("streaming")) + ); +} + +#[test] +fn declines_top_k_because_python_routes_it_by_base_model() { + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + json!({"topK": 40}) + ), + Some(Unsupported("unrecognized request parameter")) + ); +} + +#[test] +fn declines_tools_and_other_params_outside_the_allowlist() { + for param in [ + json!({"tools": []}), + json!({"tool_choice": {"auto": {}}}), + json!({"thinking": {"type": "enabled"}}), + json!({"requestMetadata": {"k": "v"}}), + json!({"outputConfig": {}}), + json!({"_parallel_tool_use_config": {}}), + ] { + assert_eq!( + reason(json!([{"role": "user", "content": "hi"}]), param.clone()), + Some(Unsupported("unrecognized request parameter")), + "expected {param} to decline" + ); + } +} + +#[test] +fn declines_blank_text_rather_than_substituting_the_anthropic_placeholder() { + for content in [ + json!(""), + json!(" "), + json!([{"type": "text", "text": " "}]), + ] { + assert_eq!( + reason( + json!([{"role": "user", "content": content}, {"role": "user", "content": "hi"}]), + json!({}) + ), + Some(Unsupported("blank message text")), + "expected blank content {content} to decline" + ); + } +} + +#[test] +fn declines_a_message_whose_content_list_is_empty() { + // The blank-text check scans parts, so an empty list clears it; Converse + // rejects an empty `content` array, which is a decline the core owes the + // host before the call rather than an error after it. + assert_eq!( + reason(json!([{"role": "user", "content": []}]), json!({})), + Some(Unsupported("message without content")) + ); + assert_eq!( + reason( + json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]), + json!({}) + ), + None + ); +} + +#[test] +fn declines_a_conversation_that_opens_or_closes_on_an_assistant_turn() { + assert_eq!( + reason( + json!([ + {"role": "assistant", "content": "prefill"}, + {"role": "user", "content": "hi"} + ]), + json!({}) + ), + Some(Unsupported( + "conversation does not run user turn to user turn" + )) + ); + assert_eq!( + reason( + json!([ + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "prefill"} + ]), + json!({}) + ), + Some(Unsupported( + "conversation does not run user turn to user turn" + )) + ); +} + +#[test] +fn accepts_a_user_to_user_text_conversation() { + assert_eq!( + reason( + json!([ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + {"role": "assistant", "content": "hello"}, + {"role": "user", "content": "again"} + ]), + json!({"maxTokens": 16}) + ), + None + ); +} + +#[test] +fn builds_the_converse_url_from_the_region_in_the_model_id() { + let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG; + assert_eq!( + config + .complete_url(None, "us-east-1/anthropic.claude-v2", &Map::new(), &|_| { + None + }) + .expect("url builds"), + "https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse" + ); +} + +#[test] +fn falls_back_to_the_region_env_then_the_default_region() { + let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG; + let with_env = |key: &str| (key == "AWS_REGION_NAME").then(|| "eu-west-1".to_string()); + assert_eq!( + config + .complete_url(None, "anthropic.claude-v2", &Map::new(), &with_env) + .expect("url builds"), + "https://bedrock-runtime.eu-west-1.amazonaws.com/model/anthropic.claude-v2/converse" + ); + assert_eq!( + config + .complete_url(None, "anthropic.claude-v2", &Map::new(), &|_| None) + .expect("url builds"), + "https://bedrock-runtime.us-west-2.amazonaws.com/model/anthropic.claude-v2/converse" + ); +} + +#[test] +fn prefers_an_explicit_runtime_endpoint_over_the_api_base() { + let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG; + let overrides = params(json!({"aws_bedrock_runtime_endpoint": "https://vpce.internal/"})); + assert_eq!( + config + .complete_url( + Some("https://ignored.example"), + "anthropic.claude-v2", + &overrides, + &|_| None + ) + .expect("url builds"), + "https://vpce.internal/model/anthropic.claude-v2/converse" + ); +} + +#[test] +fn signs_with_sigv4_in_the_resolved_region() { + let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG; + assert_eq!( + config + .auth( + None, + "eu-central-1/anthropic.claude-v2", + &Map::new(), + &|_| None + ) + .expect("auth resolves"), + ChatCompletionsAuth::AwsSigV4 { + region: "eu-central-1".to_string() + } + ); +} + +#[test] +fn a_bearer_token_outranks_sigv4_the_way_python_resolves_it() { + // Python's get_request_headers reads `api_key` as the Bedrock bearer token + // and only falls back to the env when the caller passed none, so each case + // pins one of its precedence rules. Signing as the host principal when a + // bearer identity is configured would cross an account and quota boundary. + let bedrock_env = + |key: &str| (key == "AWS_BEARER_TOKEN_BEDROCK").then(|| "from-env".to_string()); + let no_env = |_: &str| None; + let resolve = |api_key, env: &dyn Fn(&str) -> Option| { + BEDROCK_CHAT_COMPLETIONS_CONFIG + .auth( + api_key, + "eu-central-1/anthropic.claude-v2", + &Map::new(), + env, + ) + .expect("auth resolves") + }; + let bearer = |token: &str| ChatCompletionsAuth::Bearer { + token: token.to_string(), + }; + let sigv4 = ChatCompletionsAuth::AwsSigV4 { + region: "eu-central-1".to_string(), + }; + + // A caller-supplied key is the bearer token, and outranks the env. + assert_eq!( + resolve(Some("bedrock-api-key"), &bedrock_env), + bearer("bedrock-api-key") + ); + // No key, so the env supplies it. + assert_eq!(resolve(None, &bedrock_env), bearer("from-env")); + // An empty key is not a bearer token, and deliberately does NOT reach for + // the env, which is what Python's `is not None` check does. + assert_eq!(resolve(Some(""), &bedrock_env), sigv4); + // Whitespace is truthy in Python, so it stays a bearer token rather than + // silently becoming a host-credentialed SigV4 request. + assert_eq!(resolve(Some(" "), &no_env), bearer(" ")); + // Neither present, so SigV4 as before. + assert_eq!(resolve(None, &no_env), sigv4); +} + +#[test] +fn normalizes_a_converse_response_into_openai_shape() { + let response = transform_response(json!({ + "output": {"message": {"role": "assistant", "content": [ + {"text": "hello"}, {"text": " there"} + ]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15} + })) + .expect("response transforms"); + + assert_eq!(response.model, "anthropic.claude-sonnet-4-5-v1:0"); + assert_eq!( + response.choices[0].message.content.as_deref(), + Some("hello there") + ); + assert_eq!(response.choices[0].finish_reason, "stop"); + assert_eq!(response.usage.prompt_tokens, 11); + assert_eq!(response.usage.completion_tokens, 4); + assert_eq!(response.usage.total_tokens, 15); +} + +#[test] +fn maps_converse_stop_reasons_python_maps() { + for (provider_reason, expected) in [ + ("end_turn", "stop"), + ("stop_sequence", "stop"), + ("max_tokens", "length"), + ("guardrail_intervened", "content_filter"), + // Converse emits this one, and Python's `_FINISH_REASON_MAP` carries + // it. Folding it into `stop` reports a filtered completion as a normal + // one to anything keying on the finish reason. + ("content_filtered", "content_filter"), + ("content_filter", "content_filter"), + ] { + let response = transform_response(json!({ + "output": {"message": {"content": [{"text": "x"}]}}, + "stopReason": provider_reason, + "usage": {"inputTokens": 1, "outputTokens": 1} + })) + .expect("response transforms"); + assert_eq!( + response.choices[0].finish_reason, expected, + "stopReason {provider_reason}" + ); + } +} + +#[test] +fn reports_an_empty_converse_answer_as_an_empty_string_not_null() { + // Converse assigns the joined text unconditionally + // (`chat_completion_message["content"] = content_str`), unlike Anthropic's + // `merged_text or None`, so an empty answer is `""` on both paths. A caller + // calling `.strip()` on it would break on the Rust path alone. Reachable + // through a filtered or guardrail-intervened response. + for content in [json!([]), json!([{"text": ""}])] { + let response = transform_response(json!({ + "output": {"message": {"content": content}}, + "stopReason": "content_filtered", + "usage": {"inputTokens": 1, "outputTokens": 0} + })) + .expect("response transforms"); + assert_eq!(response.choices[0].message.content, Some(String::new())); + } +} + +#[test] +fn reports_the_total_tokens_converse_sent_rather_than_recomputing_them() { + // Python reads `usage["totalTokens"]` straight through here, where Anthropic + // has no such field and adds the two counts instead. The two agree while the + // gate declines every cache_control request, so this is what keeps them + // agreeing if that ever widens. + let response = transform_response(json!({ + "output": {"message": {"content": [{"text": "x"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 10, "outputTokens": 4, "cacheReadInputTokens": 7, "totalTokens": 14} + })) + .expect("response transforms"); + assert_eq!( + response.usage.total_tokens, 14, + "provider total was recomputed" + ); + assert_eq!(response.usage.prompt_tokens, 17); + assert_eq!(response.usage.completion_tokens, 4); +} + +#[test] +fn falls_back_to_the_computed_total_when_converse_omits_it() { + // Python raises a KeyError on a body with no `totalTokens`. Reporting a zero + // instead would be a worse divergence than the one above, so the computed + // total stands in. + let response = transform_response(json!({ + "output": {"message": {"content": [{"text": "x"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 10, "outputTokens": 4} + })) + .expect("response transforms"); + assert_eq!(response.usage.total_tokens, 14); +} + +#[test] +fn declines_a_cache_control_message_so_widening_the_gate_is_a_red_test() { + // Converse only reports cache token counts when the request carries a + // cachePoint block, which is why the provider total and the computed one + // cannot disagree today. This is the tripwire: whoever widens the gate to + // admit prompt caching has to come back and re-check the usage mapping + // rather than discovering a silent number change in production. + assert_eq!( + reason( + json!([{"role": "user", "content": [ + {"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}} + ]}]), + json!({}) + ), + Some(Unsupported("non-text message content")) + ); +} + +#[test] +fn folds_converse_cache_tokens_into_prompt_tokens() { + let response = transform_response(json!({ + "output": {"message": {"content": [{"text": "x"}]}}, + "stopReason": "end_turn", + "usage": { + "inputTokens": 10, + "outputTokens": 2, + "cacheReadInputTokens": 5, + "cacheWriteInputTokens": 3 + } + })) + .expect("response transforms"); + assert_eq!(response.usage.prompt_tokens, 18); + assert_eq!(response.usage.prompt_tokens_details.cached_tokens, 5); + assert_eq!( + response.usage.prompt_tokens_details.cache_creation_tokens, + 3 + ); + assert_eq!(response.usage.prompt_tokens_details.text_tokens, 10); +} + +#[test] +fn declines_a_response_carrying_a_tool_use_block() { + let err = transform_response(json!({ + "output": {"message": {"content": [ + {"toolUse": {"toolUseId": "t1", "name": "f", "input": {}}} + ]}}, + "stopReason": "tool_use", + "usage": {"inputTokens": 1, "outputTokens": 1} + })) + .expect_err("tool use block"); + assert_eq!( + err, + CoreError::Unsupported("non-text response content block") + ); +} + +#[test] +fn errors_on_a_response_missing_required_fields() { + assert_eq!( + transform_response(json!("nope")).expect_err("not an object"), + CoreError::InvalidResponse("converse response is not an object".to_string()) + ); + assert_eq!( + transform_response(json!({"usage": {}})).expect_err("no output"), + CoreError::MissingField("output.message.content") + ); + assert_eq!( + transform_response(json!({"output": {"message": {"content": []}}})).expect_err("no usage"), + CoreError::MissingField("usage") + ); +} + +#[test] +fn accepts_aws_call_configuration_without_serializing_it() { + let call_config = json!({ + "maxTokens": 16, + "aws_access_key_id": "AKIA", + "aws_secret_access_key": "secret", + "aws_session_token": "token", + "aws_region_name": "us-east-1", + "aws_profile_name": "litellm-stage", + "aws_role_name": "role", + "aws_session_name": "session", + "aws_web_identity_token": "wit", + "aws_sts_endpoint": "https://sts.example", + "aws_external_id": "ext", + "aws_bedrock_runtime_endpoint": "https://vpce.internal" + }); + assert_eq!( + reason( + json!([{"role": "user", "content": "hi"}]), + call_config.clone() + ), + None + ); + let body = transform(json!([{"role": "user", "content": "hi"}]), call_config); + assert_eq!( + body, + json!({ + "inferenceConfig": {"maxTokens": 16}, + "messages": [{"role": "user", "content": [{"text": "hi"}]}] + }), + "aws call configuration must not reach the Converse body" + ); +} + +#[test] +fn leaves_a_complete_converse_url_untouched() { + let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG; + let already_built = + "https://bedrock-runtime.us-east-1.amazonaws.com/model/us.anthropic.claude-v2%3A0/converse"; + assert_eq!( + config + .complete_url( + Some(already_built), + "anthropic.claude-v2", + &Map::new(), + &|_| None + ) + .expect("url builds"), + already_built, + "a host that encoded the model id itself must not have it re-derived" + ); +} + +#[test] +fn host_supplied_credentials_outrank_ambient_profile_and_role_state() { + use crate::providers::bedrock::aws_base::host_supplied_credentials; + + let supplied = params(json!({ + "aws_access_key_id": "AKIAHOST", + "aws_secret_access_key": "hostsecret", + "aws_session_token": "hosttoken" + })); + let credentials = host_supplied_credentials(&supplied).expect("host credentials"); + assert_eq!(credentials.access_key_id(), "AKIAHOST"); + assert_eq!(credentials.secret_access_key(), "hostsecret"); + assert_eq!(credentials.session_token(), Some("hosttoken")); + + // Without a full static pair there is nothing to honor, so the core falls + // back to deriving credentials itself. + assert!(host_supplied_credentials(¶ms(json!({"aws_access_key_id": "AKIA"}))).is_none()); + assert!( + host_supplied_credentials(¶ms( + json!({"aws_access_key_id": " ", "aws_secret_access_key": "s"}) + )) + .is_none() + ); + assert!(host_supplied_credentials(&Map::new()).is_none()); +} diff --git a/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs new file mode 100644 index 00000000000..b107950748e --- /dev/null +++ b/litellm-rust/crates/core/src/providers/bedrock/chat_completions/transformation.rs @@ -0,0 +1,297 @@ +use serde_json::{Map, Value, json}; + +use crate::chat_completions::conversation::{Conversation, TurnRole, build_conversation}; +use crate::chat_completions::response_utils::{finish_reason_for, unix_now, usage_from_parts}; +use crate::chat_completions::transformation::{ + ChatCompletionsAuth, ChatCompletionsProviderConfig, Unsupported, unsupported_message, + unsupported_param, +}; +use crate::chat_completions::types::{ + ChatCompletionsChoice, ChatCompletionsChoiceMessage, ChatCompletionsResponse, + ChatCompletionsUsage, ChatMessage, ChatMessageContent, ProviderChatRequestData, + ProviderChatResponseData, +}; +use crate::error::{CoreError, CoreResult}; + +use super::super::aws_base::{bedrock_model_id_and_region, resolve_bedrock_region}; +use super::super::constants::{AWS_BEARER_TOKEN_BEDROCK, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE}; + +/// Converse parameter names, post `map_openai_params`, that the Rust path can +/// place verbatim in `inferenceConfig`. +/// +/// `topK` is deliberately absent: Python routes it to +/// `additionalModelRequestFields` for Anthropic base models and to +/// `inferenceConfig` otherwise, and that branch reads the model catalog the +/// core cannot see. +const SUPPORTED_PARAMS: &[&str] = &["maxTokens", "temperature", "topP", "stopSequences"]; + +/// Params that belong in `inferenceConfig`, in the order Python's +/// `AmazonConverseConfig` declares them, so bodies compare cleanly. +const INFERENCE_CONFIG_PARAMS: &[&str] = SUPPORTED_PARAMS; + +const AWS_BEDROCK_RUNTIME_ENDPOINT: &str = "aws_bedrock_runtime_endpoint"; + +/// AWS call configuration a host passes down: consumed for signing and endpoint +/// resolution, never serialized into the Converse body. +const CONFIG_PARAMS: &[&str] = &[ + "aws_access_key_id", + "aws_secret_access_key", + "aws_session_token", + "aws_region_name", + "aws_session_name", + "aws_profile_name", + "aws_role_name", + "aws_web_identity_token", + "aws_sts_endpoint", + "aws_external_id", + AWS_BEDROCK_RUNTIME_ENDPOINT, +]; + +const CONVERSE_PATH_SUFFIX: &str = "/converse"; + +pub struct BedrockChatCompletionsConfig; + +pub const BEDROCK_CHAT_COMPLETIONS_CONFIG: BedrockChatCompletionsConfig = + BedrockChatCompletionsConfig; + +fn converse_body(conversation: &Conversation, params: &Map) -> Value { + let messages: Vec = conversation + .turns + .iter() + .map(|turn| { + json!({ + "role": turn.role.as_str(), + "content": turn.texts.iter().map(|text| json!({"text": text})).collect::>(), + }) + }) + .collect(); + + let inference_config = Map::from_iter(INFERENCE_CONFIG_PARAMS.iter().filter_map(|name| { + params + .get(*name) + .map(|value| ((*name).to_string(), value.clone())) + })); + + let system: Vec = conversation + .system + .iter() + .map(|text| json!({"text": text})) + .collect(); + + Value::Object(Map::from_iter( + [ + ( + "inferenceConfig".to_string(), + Value::Object(inference_config), + ), + ("messages".to_string(), json!(messages)), + ] + .into_iter() + .chain((!system.is_empty()).then(|| ("system".to_string(), json!(system)))), + )) +} + +fn has_blank_text(message: &ChatMessage) -> bool { + match &message.content { + None => false, + Some(ChatMessageContent::Text(text)) => text.trim().is_empty(), + Some(ChatMessageContent::Parts(parts)) => parts.iter().any(|part| { + part.get("text") + .and_then(Value::as_str) + .is_none_or(|text| text.trim().is_empty()) + }), + } +} + +impl ChatCompletionsProviderConfig for BedrockChatCompletionsConfig { + fn complete_url( + &self, + api_base: Option<&str>, + model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + let (model_id, model_region) = bedrock_model_id_and_region(model); + let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup); + let endpoint = optional_params + .get(AWS_BEDROCK_RUNTIME_ENDPOINT) + .and_then(Value::as_str) + .or(api_base) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .unwrap_or_else(|| BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", ®ion)); + let endpoint = endpoint.trim_end_matches('/'); + // A host that already built the full Converse URL (LiteLLM's Python + // path encodes the model id itself) passes it through untouched, the + // way the Anthropic config leaves a complete `/v1/messages` URL alone. + if endpoint.ends_with(CONVERSE_PATH_SUFFIX) { + return Ok(endpoint.to_string()); + } + Ok(format!("{endpoint}/model/{model_id}{CONVERSE_PATH_SUFFIX}")) + } + + fn auth( + &self, + api_key: Option<&str>, + model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + // Python reads `api_key` as the Bedrock bearer token and consults the + // env only when the caller passed none, so a caller-supplied empty key + // falls through to SigV4 without reaching for the environment. An + // all-whitespace token stays a bearer token here because Python sends + // it too: treating it as absent would sign as the host principal + // instead, which is the identity swap this branch exists to prevent. + let bearer = match api_key { + Some(key) => Some(key.to_string()), + None => env_lookup(AWS_BEARER_TOKEN_BEDROCK), + } + .filter(|token| !token.is_empty()); + if let Some(token) = bearer { + return Ok(ChatCompletionsAuth::Bearer { token }); + } + let (_, model_region) = bedrock_model_id_and_region(model); + Ok(ChatCompletionsAuth::AwsSigV4 { + region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup), + }) + } + + fn default_headers(&self) -> &'static [(&'static str, &'static str)] { + &[("Content-Type", "application/json")] + } + + fn supported_params(&self) -> &'static [&'static str] { + SUPPORTED_PARAMS + } + + fn config_params(&self) -> &'static [&'static str] { + CONFIG_PARAMS + } + + fn unsupported_reason( + &self, + messages: &[ChatMessage], + optional_params: &Map, + ) -> Option { + unsupported_param(SUPPORTED_PARAMS, CONFIG_PARAMS, optional_params) + .or_else(|| messages.iter().find_map(unsupported_message)) + // Python's Converse translation drops blank text blocks instead of + // substituting the placeholder the shared conversation builder + // applies, so decline blank text rather than diverge. + .or_else(|| { + messages + .iter() + .any(has_blank_text) + .then_some(Unsupported("blank message text")) + }) + // Converse has no assistant prefill: Python inserts a continue turn + // when a conversation opens or closes on an assistant message, and + // only under `litellm.modify_params`, which the core cannot see. + // Declining both ends also keeps the shared builder's final + // assistant right-strip (an Anthropic rule) unreachable here. + .or_else(|| { + let conversation = build_conversation(messages); + let ends_on_assistant = conversation + .turns + .last() + .is_some_and(|turn| turn.role == TurnRole::Assistant); + (!conversation.opens_on_user_turn() || ends_on_assistant).then_some(Unsupported( + "conversation does not run user turn to user turn", + )) + }) + } + + fn transform_request( + &self, + _model: &str, + messages: Vec, + optional_params: Map, + ) -> CoreResult { + Ok(ProviderChatRequestData { + body: converse_body(&build_conversation(&messages), &optional_params), + }) + } + + fn transform_response( + &self, + model: &str, + response: ProviderChatResponseData, + ) -> CoreResult { + let body = response.body.as_object().ok_or_else(|| { + CoreError::InvalidResponse("converse response is not an object".into()) + })?; + + let content = body + .get("output") + .and_then(|output| output.get("message")) + .and_then(|message| message.get("content")) + .and_then(Value::as_array) + .ok_or(CoreError::MissingField("output.message.content"))?; + // The route declines tool requests, so anything other than a text block + // is something this path never asked for. Decline; the host falls back. + if content.iter().any(|block| { + block + .as_object() + .is_none_or(|block| block.len() != 1 || !block.contains_key("text")) + }) { + return Err(CoreError::Unsupported("non-text response content block")); + } + let text: String = content + .iter() + .filter_map(|block| block.get("text").and_then(Value::as_str)) + .collect(); + + let usage = body + .get("usage") + .and_then(Value::as_object) + .ok_or(CoreError::MissingField("usage"))?; + let field = |name: &str| usage.get(name).and_then(Value::as_u64).unwrap_or(0); + let computed = usage_from_parts( + field("inputTokens"), + field("outputTokens"), + field("cacheReadInputTokens"), + field("cacheWriteInputTokens"), + ); + // Converse reports `totalTokens` and Python passes it straight through, + // where Anthropic has no such field and Python adds the two counts + // instead, so only this provider overrides the computed total. Python + // does a bare `usage["totalTokens"]` lookup, so a body without the key + // raises there rather than reporting a zero; fall back to the computed + // total, which is the closest thing to that without failing the call. + let usage = ChatCompletionsUsage { + total_tokens: usage + .get("totalTokens") + .and_then(Value::as_u64) + .unwrap_or(computed.total_tokens), + ..computed + }; + + Ok(ChatCompletionsResponse { + created: unix_now(), + // Converse echoes no model id, so Python reports the requested one. + model: model.to_string(), + choices: vec![ChatCompletionsChoice { + index: 0, + message: ChatCompletionsChoiceMessage { + role: "assistant".to_string(), + // Converse assigns the joined string unconditionally, so an + // empty response is `""` here and not `None` as it is on + // Anthropic. A caller calling `.strip()` on it would break + // on this path alone. + content: Some(text), + }, + finish_reason: finish_reason_for( + body.get("stopReason").and_then(Value::as_str).unwrap_or(""), + ) + .to_string(), + }], + usage, + }) + } +} + +#[cfg(test)] +#[path = "tests.rs"] +mod tests; diff --git a/litellm-rust/crates/core/src/providers/bedrock/constants.rs b/litellm-rust/crates/core/src/providers/bedrock/constants.rs index 785295207e7..be215cc9016 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/constants.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/constants.rs @@ -11,6 +11,31 @@ pub const AWS_ROLE_ARN: &str = "AWS_ROLE_ARN"; pub const AWS_WEB_IDENTITY_TOKEN_FILE: &str = "AWS_WEB_IDENTITY_TOKEN_FILE"; pub const AWS_STS_ENDPOINT: &str = "AWS_STS_ENDPOINT"; pub const AWS_EXTERNAL_ID: &str = "AWS_EXTERNAL_ID"; +pub const AWS_BEARER_TOKEN_BEDROCK: &str = "AWS_BEARER_TOKEN_BEDROCK"; + +/// Headers SigV4 covers, beyond the `x-amz-` / `x-amzn-` prefixes. Mirrors +/// Python's `_filter_headers_for_aws_signature` allowlist. +pub const AWS_SIGNED_HEADER_NAMES: &[&str] = &[ + "host", + "content-type", + "date", + "x-amz-date", + "x-amz-security-token", + "x-amz-content-sha256", + "x-amz-algorithm", + "x-amz-credential", + "x-amz-signedheaders", + "x-amz-signature", +]; +/// Headers the signer emits itself. Mirrors Python's `SIGV4_COMPUTED_HEADERS`, +/// which the reattach loop skips so a caller's copy cannot ride alongside the +/// computed one. +pub const SIGV4_COMPUTED_HEADER_NAMES: &[&str] = &[ + "authorization", + "x-amz-date", + "x-amz-security-token", + "date", +]; pub const BEDROCK_SERVICE: &str = "bedrock"; pub const DEFAULT_SESSION_NAME_PREFIX: &str = "litellm-session"; pub const DEFAULT_BEDROCK_REGION: &str = "us-west-2"; diff --git a/litellm-rust/crates/core/src/providers/bedrock/mod.rs b/litellm-rust/crates/core/src/providers/bedrock/mod.rs index b09675ad7dd..d9cd3efcb74 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/mod.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/mod.rs @@ -5,4 +5,5 @@ #[cfg(feature = "bedrock-auth")] pub mod audio_transcription; pub mod aws_base; +pub mod chat_completions; mod constants; diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index f0cc26a0cca..c6f81cf6916 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -6,6 +6,10 @@ use litellm_ai_gateway::io::audio_transcription::{ }; use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; +use litellm_core::chat_completions::types::{ChatCompletionsRequest, ChatCompletionsResponse}; +use litellm_core::chat_completions::{ + chat_completions as run_chat_completions, chat_completions_decline_reason, +}; use litellm_core::error::CoreError; use litellm_core::messages::messages as run_messages; use litellm_core::messages::types::{AnthropicMessagesResponse, MessagesRequest}; @@ -16,6 +20,20 @@ use serde_json::{Map, Value}; mod gil; +pyo3::create_exception!( + _native, + RustBridgeDeclined, + pyo3::exceptions::PyException, + "The route declined before calling the provider, so the host may retry on its own path." +); + +pyo3::create_exception!( + _native, + RustUpstreamError, + pyo3::exceptions::PyException, + "The provider call was already issued and failed. Args are (status, message); status is 0 when there was no HTTP response." +); + type MarshaledOcrInputs = ( Value, Option>, @@ -45,6 +63,15 @@ fn messages_response_to_py( json_to_py(py, value) } +fn chat_completions_response_to_py( + py: Python<'_>, + response: ChatCompletionsResponse, +) -> PyResult> { + let value = + serde_json::to_value(response).map_err(|err| PyValueError::new_err(err.to_string()))?; + json_to_py(py, value) +} + fn core_error_to_pyerr(err: CoreError) -> PyErr { match err { CoreError::Auth(message) => PyValueError::new_err(message), @@ -56,6 +83,33 @@ fn core_error_to_pyerr(err: CoreError) -> PyErr { } } +/// Map a core error for a route whose host keeps a Python implementation. +/// +/// The distinction the host needs is whether the provider was already called. +/// Everything raised before the request goes out is safe for the host to retry +/// on its own path; anything after it is not, because the provider has already +/// done the work and billed for it. +fn chat_completions_error_to_pyerr(err: CoreError) -> PyErr { + match err { + CoreError::Unsupported(_) + | CoreError::Auth(_) + | CoreError::InvalidProvider(_) + | CoreError::InvalidRequest(_) + | CoreError::InvalidType { .. } + | CoreError::MissingField(_) + | CoreError::Routing(_) + // Nothing reached the provider, so serving it on Python cannot double + // bill and is the only way the caller gets an answer at all. + | CoreError::Connect(_) => RustBridgeDeclined::new_err(err.to_string()), + CoreError::Http { status, body } => { + RustUpstreamError::new_err((status, format!("{status}: {body}"))) + } + CoreError::Network(message) | CoreError::InvalidResponse(message) => { + RustUpstreamError::new_err((0u16, message)) + } + } +} + fn optional_object_to_map( py: Python<'_>, name: &'static str, @@ -430,6 +484,143 @@ fn amessages( }) } +type MarshaledChatCompletionsInputs = ( + Value, + Map, + Option>, + Option, +); + +fn marshal_chat_completions_inputs( + py: Python<'_>, + messages: Py, + optional_params: Option>, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult { + let messages = py_to_json(py, messages.bind(py))?; + if !messages.is_array() { + return Err(PyValueError::new_err("messages must be a list")); + } + let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; + let extra_headers = match extra_headers { + Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?), + None => None, + }; + Ok(( + messages, + optional_params, + extra_headers, + optional_timeout(timeout_seconds), + )) +} + +/// The decline reason for this request, or `None` when the Rust path accepts +/// it. Resolves no credentials and performs no I/O, so a host can ask before +/// committing to either path. +#[pyfunction] +#[pyo3(signature = (model, messages, optional_params=None, custom_llm_provider=None))] +fn chat_completions_decline( + py: Python<'_>, + model: String, + messages: Py, + optional_params: Option>, + custom_llm_provider: Option, +) -> PyResult> { + let messages = py_to_json(py, messages.bind(py))?; + let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; + Ok(chat_completions_decline_reason( + &model, + custom_llm_provider.as_deref(), + messages, + &optional_params, + ) + .map(str::to_string)) +} + +#[pyfunction] +#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] +#[allow(clippy::too_many_arguments)] +fn chat_completions( + py: Python<'_>, + model: String, + messages: Py, + optional_params: Option>, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + let (messages, optional_params, extra_headers, timeout) = marshal_chat_completions_inputs( + py, + messages, + optional_params, + extra_headers, + timeout_seconds, + )?; + + let result = gil::release_gil(py, || { + pyo3_async_runtimes::tokio::get_runtime().block_on(run_chat_completions( + ChatCompletionsRequest { + model: &model, + messages, + optional_params, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + timeout, + }, + )) + }); + + match result { + Ok(response) => chat_completions_response_to_py(py, response), + Err(err) => Err(chat_completions_error_to_pyerr(err)), + } +} + +#[pyfunction] +#[pyo3(signature = (model, messages, optional_params=None, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, timeout_seconds=None))] +#[allow(clippy::too_many_arguments)] +fn achat_completions( + py: Python<'_>, + model: String, + messages: Py, + optional_params: Option>, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + timeout_seconds: Option, +) -> PyResult> { + let (messages, optional_params, extra_headers, timeout) = marshal_chat_completions_inputs( + py, + messages, + optional_params, + extra_headers, + timeout_seconds, + )?; + + pyo3_async_runtimes::tokio::future_into_py(py, async move { + let response = run_chat_completions(ChatCompletionsRequest { + model: &model, + messages, + optional_params, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + timeout, + }) + .await + .map_err(chat_completions_error_to_pyerr)?; + + Python::attach(|py| chat_completions_response_to_py(py, response)) + }) +} + #[pyfunction] fn gil_stats(py: Python<'_>) -> PyResult> { let stats = PyDict::new(py); @@ -439,12 +630,18 @@ fn gil_stats(py: Python<'_>) -> PyResult> { #[pymodule] fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> { + let py = module.py(); module.add_function(wrap_pyfunction!(ocr, module)?)?; module.add_function(wrap_pyfunction!(aocr, module)?)?; module.add_function(wrap_pyfunction!(transcription, module)?)?; module.add_function(wrap_pyfunction!(atranscription, module)?)?; module.add_function(wrap_pyfunction!(messages, module)?)?; module.add_function(wrap_pyfunction!(amessages, module)?)?; + module.add("RustBridgeDeclined", py.get_type::())?; + module.add("RustUpstreamError", py.get_type::())?; + module.add_function(wrap_pyfunction!(chat_completions_decline, module)?)?; + module.add_function(wrap_pyfunction!(chat_completions, module)?)?; + module.add_function(wrap_pyfunction!(achat_completions, module)?)?; module.add_class::()?; module.add_function(wrap_pyfunction!(gil_stats, module)?)?; Ok(()) diff --git a/litellm/_logging.py b/litellm/_logging.py index 7d3a30c6d1a..e55c6bc40a8 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -8,6 +8,12 @@ from logging import Formatter from typing import Any, Final import litellm +from litellm.constants import ( + LITELLM_TRUNCATED_PAYLOAD_FIELD, + LITELLM_TRUNCATION_STDOUT_SAFEGUARD_NOTE, + MAX_STRING_LENGTH_STDOUT_LOG, +) +from litellm.litellm_core_utils.env_utils import get_env_int from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.litellm_core_utils.safe_json_loads import safe_json_loads from litellm.litellm_core_utils.secret_redaction import redact_string, redact_structured_value @@ -101,7 +107,7 @@ class SecretRedactionFilter(logging.Filter): # Redact exception tracebacks if record.exc_info and record.exc_info[1] is not None: try: - record.exc_text = _redact_string(self._formatter.formatException(record.exc_info)) + record.exc_text = _redact_string(record.exc_text or self._formatter.formatException(record.exc_info)) except Exception: pass @@ -116,6 +122,72 @@ class SecretRedactionFilter(logging.Filter): _secret_filter: Final = SecretRedactionFilter() +def _get_max_string_length_stdout_log() -> int: + """Read the limit per record so a value loaded later via proxy config + environment_variables is honored.""" + return get_env_int("MAX_STRING_LENGTH_STDOUT_LOG", MAX_STRING_LENGTH_STDOUT_LOG) + + +def _stdout_truncation_marker(skipped_chars: int) -> str: + return ( + f"... ({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped {skipped_chars} chars. " + f"{LITELLM_TRUNCATION_STDOUT_SAFEGUARD_NOTE}) ..." + ) + + +def _truncate_for_stdout_log(text: str, limit: int) -> str: + kept_chars: Final = limit - len(_stdout_truncation_marker(len(text))) + if kept_chars <= 0: + return text[:limit] + head_chars: Final = kept_chars // 2 + tail_chars: Final = kept_chars - head_chars + return f"{text[:head_chars]}{_stdout_truncation_marker(len(text) - kept_chars)}{text[-tail_chars:]}" + + +class StdoutLogTruncationFilter(logging.Filter): + """Bounds how much of an oversized log line reaches stdout. + + A provider error string can echo the whole request payload, so one failed agentic + request writes hundreds of KB to stdout, repeatedly as the exception propagates from + the router to the proxy handler and into its traceback, all inline on the event loop. + + DEBUG records pass through untouched, since dumping full payloads is the point of + `--detailed_debug`, and logging callbacks (OTEL, Datadog, etc.) don't run through + logging filters at all, so they still get the untruncated error. + """ + + _formatter = logging.Formatter() + + def filter(self, record: logging.LogRecord) -> bool: + if record.levelno < logging.INFO: + return True + + limit: Final = _get_max_string_length_stdout_log() + if limit <= 0: + return True + + try: + message: Final = record.getMessage() + except (TypeError, ValueError): + return True + + if len(message) > limit: + record.msg = _truncate_for_stdout_log(message, limit) # rebind-ok: the Filter interface mutates the record + record.args = None # rebind-ok: args are consumed by the truncated message above + + if isinstance(record.exc_info, tuple): + exc_text: Final = record.exc_text or self._formatter.formatException(record.exc_info) + if len(exc_text) > limit: + record.exc_text = _truncate_for_stdout_log( # rebind-ok: the Filter interface mutates the record + exc_text, limit + ) + + return True + + +_stdout_truncation_filter: Final = StdoutLogTruncationFilter() + + class CorrelationContextFilter(logging.Filter): """Stamps each log record with the current request's trace_id and session_id from contextvars. @@ -301,6 +373,7 @@ def _setup_json_exception_handlers(formatter): error_handler: Final = logging.StreamHandler() error_handler.setFormatter(formatter) error_handler.addFilter(_secret_filter) + error_handler.addFilter(_stdout_truncation_filter) error_handler.addFilter(_correlation_filter) # Setup excepthook for uncaught exceptions @@ -365,6 +438,12 @@ verbose_router_logger.addHandler(handler) verbose_proxy_logger.addHandler(handler) verbose_logger.addHandler(handler) +# Filters attached to the logger, not the handler, survive callers swapping in their own +# handlers (JSON mode, uvicorn log config, a host app's root handler). +verbose_router_logger.addFilter(_stdout_truncation_filter) +verbose_proxy_logger.addFilter(_stdout_truncation_filter) +verbose_logger.addFilter(_stdout_truncation_filter) + def _suppress_loggers(): """Suppress noisy loggers at INFO level""" diff --git a/litellm/completion_extras/litellm_responses_transformation/handler.py b/litellm/completion_extras/litellm_responses_transformation/handler.py index 33206629b41..727c39c16ec 100644 --- a/litellm/completion_extras/litellm_responses_transformation/handler.py +++ b/litellm/completion_extras/litellm_responses_transformation/handler.py @@ -25,6 +25,11 @@ class ResponsesToCompletionBridgeHandlerInputKwargs(TypedDict): encoding: object +def _restore_routing_prefix(model: str, custom_llm_provider: str) -> str: + """`responses()` runs `get_llm_provider()` itself, so hand back the prefixed model `completion()` started from.""" + return f"{custom_llm_provider}/{model}" + + class ResponsesToCompletionBridgeHandler: def __init__(self): from .transformation import LiteLLMResponsesTransformationHandler @@ -184,14 +189,11 @@ class ResponsesToCompletionBridgeHandler: client=kwargs.get("client"), ) - # Pin the resolved provider so `responses()` doesn't re-run - # `get_llm_provider()` on the model string and strip a second - # provider prefix (see GitHub issue #28505). request_data already - # carries `custom_llm_provider` via the spread of - # `sanitized_litellm_params`; overwriting it on the dict (rather - # than adding an explicit kwarg) avoids the duplicate-keyword - # TypeError that would otherwise fire on the real bridge path. + # Set on request_data rather than passed as explicit kwargs: the spread of + # `sanitized_litellm_params` already carries both, so passing them again + # would raise a duplicate-keyword TypeError. request_data["custom_llm_provider"] = custom_llm_provider + request_data["model"] = _restore_routing_prefix(model, custom_llm_provider) result: Final = responses( **request_data, ) @@ -282,13 +284,11 @@ class ResponsesToCompletionBridgeHandler: except Exception as e: raise e - # Pin the resolved provider so `aresponses()` doesn't re-run - # `get_llm_provider()` on the model string and strip a second - # provider prefix (see GitHub issue #28505). Set on request_data - # rather than passed as a separate kwarg to avoid the duplicate- - # keyword TypeError when `sanitized_litellm_params` already - # carries `custom_llm_provider`. + # Set on request_data rather than passed as explicit kwargs: the spread of + # `sanitized_litellm_params` already carries both, so passing them again + # would raise a duplicate-keyword TypeError. request_data["custom_llm_provider"] = custom_llm_provider + request_data["model"] = _restore_routing_prefix(model, custom_llm_provider) result: Final = await aresponses( **request_data, aresponses=True, diff --git a/litellm/constants.py b/litellm/constants.py index 762b7f1201c..0791721aa00 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -3,7 +3,7 @@ import sys from types import MappingProxyType from typing import Final, Literal -from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_or_none +from litellm.litellm_core_utils.env_utils import get_env_int, get_env_int_in_range, get_env_int_or_none DEFAULT_HEALTH_CHECK_PROMPT: Final = str(os.getenv("DEFAULT_HEALTH_CHECK_PROMPT", "test from litellm")) AZURE_DEFAULT_RESPONSES_API_VERSION: Final = str(os.getenv("AZURE_DEFAULT_RESPONSES_API_VERSION", "preview")) @@ -49,6 +49,8 @@ LITELLM_MAX_STREAMING_DURATION_SECONDS: Final = ( # Set to 0 to disable truncation. MAX_BASE64_LENGTH_FOR_LOGGING: Final = int(os.getenv("MAX_BASE64_LENGTH_FOR_LOGGING", 64)) +MAX_STRING_LENGTH_STDOUT_LOG: Final = get_env_int("MAX_STRING_LENGTH_STDOUT_LOG", 4096) + # When true, adds detailed per-phase timing breakdown headers to responses. # Headers: x-litellm-timing-{pre-processing,llm-api,post-processing,message-copy}-ms LITELLM_DETAILED_TIMING: Final = os.getenv("LITELLM_DETAILED_TIMING", "false").lower() == "true" @@ -323,6 +325,17 @@ DEFAULT_MOCK_RESPONSE_PROMPT_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_MOCK_RE DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT: Final = int(os.getenv("DEFAULT_MOCK_RESPONSE_COMPLETION_TOKEN_COUNT", 20)) MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES: Final = int(os.getenv("MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES", 768)) MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES: Final = int(os.getenv("MAX_LONG_SIDE_FOR_IMAGE_HIGH_RES", 2000)) +# tiktoken's BPE merge loop is quadratic in the length of a single regex piece, so a long run of one +# repeated character (dot leaders, whitespace, zero-padded base64) can take minutes on a multi-MB payload. +# Encoding in chunks makes the cost linear, at a drift of at most ~1 token per chunk boundary. The upper +# bound keeps a misconfigured chunk size from restoring the quadratic cost this exists to remove. +TIKTOKEN_ENCODE_MAX_CHUNK_SIZE_CHARS: Final = 4096 +TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS: Final = get_env_int_in_range( + "TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS", + default=1024, + minimum=1, + maximum=TIKTOKEN_ENCODE_MAX_CHUNK_SIZE_CHARS, +) MAX_TILE_WIDTH: Final = int(os.getenv("MAX_TILE_WIDTH", 512)) MAX_TILE_HEIGHT: Final = int(os.getenv("MAX_TILE_HEIGHT", 512)) OPENAI_FILE_SEARCH_COST_PER_1K_CALLS: Final = float(os.getenv("OPENAI_FILE_SEARCH_COST_PER_1K_CALLS", 2.5 / 1000)) @@ -471,6 +484,9 @@ MAX_TIME_TO_CLEAR_QUEUE: Final = float(os.getenv("MAX_TIME_TO_CLEAR_QUEUE", 5.0) LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS: Final = float( os.getenv("LOGGING_WORKER_AGGRESSIVE_CLEAR_COOLDOWN_SECONDS", 0.5) ) # Cooldown time in seconds before allowing another aggressive clear (default: 0.5s) +LOGGING_EXECUTOR_MAX_THREADS: Final = get_env_int("LOGGING_EXECUTOR_MAX_THREADS", 100) +LOGGING_EXECUTOR_MAX_PENDING_TASKS: Final = get_env_int("LOGGING_EXECUTOR_MAX_PENDING_TASKS", 10_000) +LOGGING_EXECUTOR_DROPPED_TASK_LOG_INTERVAL_SECONDS: Final = 30.0 DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE: Final = os.getenv( "DD_TRACER_STREAMING_CHUNK_YIELD_RESOURCE", "streaming.chunk.yield" ) @@ -1346,6 +1362,11 @@ LITELLM_TRUNCATION_DB_SAFEGUARD_NOTE: Final = ( "Full, untruncated data is logged to logging callbacks (OTEL, Datadog, etc.). " "To increase the truncation limit, set `MAX_STRING_LENGTH_PROMPT_IN_DB` in your env." ) +LITELLM_TRUNCATION_STDOUT_SAFEGUARD_NOTE: Final = ( + "Truncation is a stdout logging safeguard. " + "Full, untruncated data is logged to logging callbacks (OTEL, Datadog, etc.) and at DEBUG level. " + "To increase the truncation limit, set `MAX_STRING_LENGTH_STDOUT_LOG` in your env." +) ########################### LiteLLM Proxy Specific Constants ########################### ######################################################################################## diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 7bd0a847ad8..11b15a63484 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -51,6 +51,42 @@ def to_basic_auth(auth_value: str) -> str: return base64.b64encode(auth_value.encode("utf-8")).decode() +def strip_auth_scheme(auth_value: str, scheme: str) -> str: + """Return ``auth_value`` with a leading `` `` removed, or unchanged when absent. + + Callers supply both a bare credential and a complete header value, so prefixing + unconditionally yields ``Bearer Bearer ``. Scheme names are case-insensitive per + RFC 7235. A credential is required after the scheme, so both a token that merely begins + with the scheme text and a scheme with nothing behind it are returned untouched. + Surrounding whitespace is left to ``_strip_header_whitespace`` at header-build time. + """ + scheme_name, _, remainder = auth_value.lstrip().partition(" ") + credential: Final = remainder.lstrip() + if credential and scheme_name.lower() == scheme.lower(): + return credential + return auth_value + + +def to_basic_credentials(auth_value: str) -> str: + """Return the base64 credentials for a ``Basic`` header, encoding only when needed. + + ``Basic `` carries credentials that are already encoded, so encoding the whole + value again would bury the scheme inside the payload. This has to run before + :func:`to_basic_auth` rather than at header-build time, where no prefix is left to find. + A schemed value whose remainder does not decode is the bare ``username:password`` shape with + the scheme written in front of it, and is encoded rather than forwarded as an invalid header; + a pair always contains ``:``, which is outside the base64 alphabet, so the two never collide. + """ + credentials: Final = strip_auth_scheme(auth_value, "Basic") + if credentials == auth_value: + return to_basic_auth(auth_value) + try: + base64.b64decode(credentials, validate=True) + except ValueError: + return to_basic_auth(credentials) + return credentials + + def _strip_header_whitespace(headers: dict[str, str]) -> dict[str, str]: return { (key.strip() if isinstance(key, str) else key): (value.strip() if isinstance(value, str) else value) @@ -441,16 +477,15 @@ class MCPClient: except BaseException as e: verbose_logger.debug("Error during http_client cleanup: %s", e) - def update_auth_value(self, mcp_auth_value: str | dict[str, str]): + def update_auth_value(self, mcp_auth_value: str | dict[str, str]) -> None: """ Set the authentication header for the MCP client. """ if isinstance(mcp_auth_value, dict): self._mcp_auth_value = mcp_auth_value + elif self.auth_type == MCPAuth.basic: + self._mcp_auth_value = to_basic_credentials(mcp_auth_value) else: - if self.auth_type == MCPAuth.basic: - # Assuming mcp_auth_value is in format "username:password", convert it when updating - mcp_auth_value = to_basic_auth(mcp_auth_value) self._mcp_auth_value = mcp_auth_value def _get_auth_headers(self) -> dict: @@ -459,19 +494,20 @@ class MCPClient: if self._mcp_auth_value: if isinstance(self._mcp_auth_value, str): if self.auth_type == MCPAuth.bearer_token: - headers["Authorization"] = f"Bearer {self._mcp_auth_value}" + headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}" elif self.auth_type == MCPAuth.basic: headers["Authorization"] = f"Basic {self._mcp_auth_value}" elif self.auth_type == MCPAuth.api_key: headers["X-API-Key"] = self._mcp_auth_value elif self.auth_type == MCPAuth.authorization: + # This auth type means the caller owns the whole header value. headers["Authorization"] = self._mcp_auth_value elif self.auth_type == MCPAuth.oauth2: - headers["Authorization"] = f"Bearer {self._mcp_auth_value}" + headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}" elif self.auth_type == MCPAuth.token: - headers["Authorization"] = f"token {self._mcp_auth_value}" + headers["Authorization"] = f"token {strip_auth_scheme(self._mcp_auth_value, 'token')}" elif self.auth_type == MCPAuth.oauth2_token_exchange: - headers["Authorization"] = f"Bearer {self._mcp_auth_value}" + headers["Authorization"] = f"Bearer {strip_auth_scheme(self._mcp_auth_value, 'Bearer')}" elif isinstance(self._mcp_auth_value, dict): headers.update(self._mcp_auth_value) # Note: aws_sigv4 auth is not handled here — SigV4 requires per-request diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 1258c7593b4..f4f3b00dda0 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -27,11 +27,15 @@ from litellm.types.integrations.anthropic_cache_control_hook import ( CacheControlInjectionPoint, CacheControlMessageInjectionPoint, ) -from litellm.types.llms.anthropic import AnthropicSystemMessageContent +from litellm.types.llms.anthropic import ( + AllAnthropicToolsValues, + AnthropicSystemMessageContent, +) from litellm.types.llms.openai import ( AllMessageValues, ChatCompletionCachedContent, ChatCompletionTextObject, + ChatCompletionToolParam, PromptCacheBreakpoint, PromptCacheOptions, ) @@ -57,6 +61,8 @@ OPENAI_PROMPT_CACHE_BREAKPOINT_BLOCK_TYPES: Final = frozenset( OPENAI_API_HOST: Final = "api.openai.com" OPENAI_API_BASE_ENV_VARS: Final = ("OPENAI_BASE_URL", "OPENAI_API_BASE") +AllToolParamValues = ChatCompletionToolParam | AllAnthropicToolsValues + def supports_openai_prompt_cache_breakpoint(model: str) -> bool: model_map_flag: Final = _model_map_prompt_cache_breakpoint_flag(model) @@ -625,6 +631,50 @@ class AnthropicCacheControlHook(CustomPromptManagement): ] return points + @staticmethod + def messages_with_default_injections( + messages: list[AllMessageValues], + models: Iterable[str], + tools: list[AllToolParamValues] | None = None, + enable_prompt_caching: bool | None = None, + ) -> list[AllMessageValues]: + """Return the messages auto prompt caching will send, default breakpoints included. + + Router cache affinity depends on this. Deployment selection runs before the injection in + `litellm.acompletion`, so it has to reproduce the markers to derive the same cache key the + success event later writes from the sent messages. `models` is every candidate model of the + group: the first that would auto-inject decides, since the default breakpoints (system + prompt and trailing turn) do not depend on which deployment serves the call. Returns the + input list itself when auto-injection would not apply + """ + points: Final = next( + ( + candidate + for candidate in ( + AnthropicCacheControlHook.get_default_injection_points( + messages=messages, + system=None, + model=model, + custom_llm_provider=None, + tools=tools, + enable_prompt_caching=enable_prompt_caching, + ) + for model in models + ) + if candidate + ), + None, + ) + if not points: + return messages + return AnthropicCacheControlHook._apply_message_injections( + points=cast( # cast-ok: the default points are all message-location points + list[CacheControlMessageInjectionPoint], points + ), + messages=copy.deepcopy(messages), + max_blocks=MAX_CACHE_CONTROL_BLOCKS, + ) + @staticmethod def maybe_seed_default_injection_points( non_default_params: dict[str, Any], diff --git a/litellm/integrations/datadog/datadog_cost_management.py b/litellm/integrations/datadog/datadog_cost_management.py index b30700e98f2..7255c9c761c 100644 --- a/litellm/integrations/datadog/datadog_cost_management.py +++ b/litellm/integrations/datadog/datadog_cost_management.py @@ -11,6 +11,7 @@ from litellm.integrations.datadog.datadog_handler import ( get_datadog_hostname, get_datadog_pod_name, get_datadog_service, + normalize_datadog_tag_value, ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.custom_httpx.http_handler import ( @@ -184,7 +185,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): # Backwards-compat: team/user/model_group preserved regardless of allowlist. if metadata.get("user_api_key_alias"): - tags["user"] = str(metadata["user_api_key_alias"]) + tags["user"] = normalize_datadog_tag_value(metadata["user_api_key_alias"]) team_tag: Final = ( metadata.get("user_api_key_team_alias") or metadata.get("team_alias") @@ -192,7 +193,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): or metadata.get("team_id") ) if team_tag: - tags["team"] = str(team_tag) + tags["team"] = normalize_datadog_tag_value(team_tag) if metadata.get("model_group"): tags["model_group"] = str(metadata["model_group"]) @@ -229,7 +230,7 @@ class DatadogCostManagementLogger(CustomBatchLogger): value, ) return - tags[key] = value + tags[key] = normalize_datadog_tag_value(value) @staticmethod def _add_tag(tags: dict[str, str], key: str, value: Any) -> None: diff --git a/litellm/integrations/datadog/datadog_handler.py b/litellm/integrations/datadog/datadog_handler.py index 2450382a192..d360dac121c 100644 --- a/litellm/integrations/datadog/datadog_handler.py +++ b/litellm/integrations/datadog/datadog_handler.py @@ -3,6 +3,7 @@ from __future__ import annotations import os +import re from typing import Final from litellm.types.utils import StandardLoggingPayload @@ -36,6 +37,13 @@ def get_datadog_pod_name() -> str: return os.getenv("POD_NAME", "unknown") +def normalize_datadog_tag_value(value: object) -> str: + normalized_value: Final = "".join( + character if character.isalnum() or character in "_-:./" else "_" for character in str(value).lower() + ) + return re.sub(r"_+", "_", normalized_value).strip("_") + + def get_datadog_tags( standard_logging_object: StandardLoggingPayload | None = None, ) -> list[str]: @@ -58,7 +66,7 @@ def get_datadog_tags( if standard_logging_object: request_tags: Final = standard_logging_object.get("request_tags", []) or [] - tags.extend(f"request_tag:{tag}" for tag in request_tags) + tags.extend(f"request_tag:{normalize_datadog_tag_value(tag)}" for tag in request_tags) # Add Team Tag metadata: Final = standard_logging_object.get("metadata", {}) or {} @@ -69,6 +77,6 @@ def get_datadog_tags( or metadata.get("team_id") ) if team_tag: - tags.append(f"team:{team_tag}") + tags.append(f"team:{normalize_datadog_tag_value(team_tag)}") return tags diff --git a/litellm/integrations/datadog/datadog_metrics.py b/litellm/integrations/datadog/datadog_metrics.py index 89f990cf661..5dda336dc94 100644 --- a/litellm/integrations/datadog/datadog_metrics.py +++ b/litellm/integrations/datadog/datadog_metrics.py @@ -12,6 +12,7 @@ from litellm.integrations.datadog.datadog_handler import ( get_datadog_hostname, get_datadog_pod_name, get_datadog_service, + normalize_datadog_tag_value, ) from litellm.litellm_core_utils.safe_json_dumps import safe_dumps from litellm.llms.custom_httpx.http_handler import ( @@ -97,7 +98,7 @@ class DatadogMetricsLogger(CustomBatchLogger): ) if team_tag: - tags.append(f"team:{team_tag}") + tags.append(f"team:{normalize_datadog_tag_value(team_tag)}") return tags diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 6df04ff622d..76066f4a305 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -4067,9 +4067,10 @@ class PrometheusLogger(CustomLogger): require_auth (bool, optional): Whether to require authentication for the metrics endpoint. Defaults to False. """ - from prometheus_client import make_asgi_app + from prometheus_client import REGISTRY from litellm._logging import verbose_proxy_logger + from litellm.integrations.prometheus_metrics_endpoint import make_metrics_asgi_app from litellm.proxy.proxy_server import app # Create metrics ASGI app @@ -4078,9 +4079,9 @@ class PrometheusLogger(CustomLogger): registry: Final = CollectorRegistry() multiprocess.MultiProcessCollector(registry) - metrics_app = make_asgi_app(registry) + metrics_app = make_metrics_asgi_app(registry) else: - metrics_app = make_asgi_app() + metrics_app = make_metrics_asgi_app(REGISTRY) # Mount the metrics app to the app app.mount("/metrics", metrics_app) diff --git a/litellm/integrations/prometheus_metrics_endpoint.py b/litellm/integrations/prometheus_metrics_endpoint.py new file mode 100644 index 00000000000..b41cc13a04f --- /dev/null +++ b/litellm/integrations/prometheus_metrics_endpoint.py @@ -0,0 +1,100 @@ +"""ASGI app for `/metrics` that keeps registry rendering off the event loop. + +``prometheus_client.make_asgi_app`` collects and serializes the whole registry +inline in the coroutine, so a large scrape (tens of MB on high cardinality +deployments) blocks every other request on the loop for its whole duration. This +app renders in a worker thread instead, shares one render across concurrent +scrapes that want the same output, and streams the payload back in chunks. +""" + +from __future__ import annotations + +import asyncio +import gzip +from collections.abc import Callable, Iterator, Mapping +from dataclasses import dataclass +from types import MappingProxyType +from typing import Final + +from prometheus_client import CollectorRegistry +from prometheus_client.exposition import choose_encoder, gzip_accepted +from starlette.requests import Request +from starlette.responses import StreamingResponse +from starlette.types import ASGIApp, Receive, Scope, Send + +RESPONSE_CHUNK_SIZE_BYTES: Final = 64 * 1024 + +_GZIP_HEADERS: Final = MappingProxyType({"Content-Encoding": "gzip"}) + + +@dataclass(frozen=True, slots=True) +class ScrapeRequest: + """What a scrape asks for, normalized so that header spellings sharing an output share a render.""" + + encoder: Callable[[CollectorRegistry], bytes] + content_type: str + gzipped: bool + metric_names: tuple[str, ...] + + +def parse_scrape_request(accept: str, accept_encoding: str, metric_names: tuple[str, ...]) -> ScrapeRequest: + encoder, content_type = choose_encoder(accept) + return ScrapeRequest( + encoder=encoder, + content_type=content_type, + gzipped=gzip_accepted(accept_encoding), + metric_names=metric_names, + ) + + +def render_scrape(registry: CollectorRegistry, request: ScrapeRequest) -> bytes: + rendered: Final = request.encoder( + registry.restricted_registry(request.metric_names) if request.metric_names else registry # pyright: ignore[reportArgumentType] # RestrictedRegistry is registry-shaped but not a subclass + ) + return gzip.compress(rendered) if request.gzipped else rendered + + +class CoalescedScrapeRenderer: + """Renders the registry in a worker thread, sharing one render per distinct output across concurrent scrapes.""" + + def __init__(self, registry: CollectorRegistry) -> None: + self._registry = registry + self._inflight: Mapping[ScrapeRequest, asyncio.Task[bytes]] = MappingProxyType({}) + + def _forget(self, finished: asyncio.Task[bytes]) -> None: + self._inflight = MappingProxyType({key: task for key, task in self._inflight.items() if task is not finished}) + + async def render(self, request: ScrapeRequest) -> bytes: + inflight: Final = self._inflight.get(request) + if inflight is not None: + return await asyncio.shield(inflight) + + task: Final = asyncio.create_task(asyncio.to_thread(render_scrape, self._registry, request)) + self._inflight = MappingProxyType({**self._inflight, request: task}) + task.add_done_callback(self._forget) + return await asyncio.shield(task) + + +def _chunks(body: bytes) -> Iterator[bytes]: + return (body[start : start + RESPONSE_CHUNK_SIZE_BYTES] for start in range(0, len(body), RESPONSE_CHUNK_SIZE_BYTES)) + + +def make_metrics_asgi_app(registry: CollectorRegistry) -> ASGIApp: + renderer: Final = CoalescedScrapeRenderer(registry) + + async def metrics_app(scope: Scope, receive: Receive, send: Send) -> None: + request: Final = Request(scope, receive) + scrape: Final = parse_scrape_request( + accept=request.headers.get("accept", ""), + accept_encoding=request.headers.get("accept-encoding", ""), + metric_names=tuple(request.query_params.getlist("name[]")), + ) + body: Final = await renderer.render(scrape) + response: Final = StreamingResponse( + _chunks(body), + media_type=scrape.content_type, + headers=_GZIP_HEADERS if scrape.gzipped else None, + ) + await response(scope, receive, send) + + return metrics_app diff --git a/litellm/integrations/shadow_eval_logger.py b/litellm/integrations/shadow_eval_logger.py index da02db4e44b..5f4e7c71395 100644 --- a/litellm/integrations/shadow_eval_logger.py +++ b/litellm/integrations/shadow_eval_logger.py @@ -10,7 +10,7 @@ import asyncio import hashlib import random import traceback -from collections.abc import Callable, Mapping, Sequence +from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timezone from itertools import groupby @@ -42,8 +42,9 @@ if TYPE_CHECKING: from litellm.router import Router from litellm.types.utils import StandardLoggingPayload -# A job starting, stopping, or hitting its turn budget propagates to sampling within one -# TTL; the turn budget can overshoot by at most one TTL of in-flight samples per pod. +# A job starting, stopping, or hitting a budget propagates to sampling within one TTL; +# the spend gate re-checks the cross-pod counter at pipeline entry, so it overshoots +# only by the samples already in flight when the cap is crossed. _JOBS_CACHE_TTL_SECONDS: Final = 10 # Concurrent shadow+judge pipelines per pod: a traffic spike turns into skipped samples @@ -340,13 +341,24 @@ def _failure_detail(e: BaseException) -> str: return f"{type(e).__name__}{location}: {e}" -def _judge_call_cost(response: object) -> float: - """Price a judge call, treating an unmapped judge model as free rather than fatal.""" +def _call_cost(response: object) -> float: + """Price one eval-arm call with the figure the spend pipeline bills: the router client + stamps _hidden_params.response_cost from the deployment's own pricing, which the public + price map lookup below cannot see (it reads 0 for deployment-priced models).""" + getter: Final = getattr(getattr(response, "_hidden_params", None), "get", None) + stamped: Final = getter("response_cost") if callable(getter) else None + if isinstance(stamped, (int, float)): + return float(stamped) + return _price_map_cost(response) + + +def _price_map_cost(response: object) -> float: + """Public price map fallback, treating an unmapped model as free rather than fatal.""" import litellm try: return litellm.completion_cost(completion_response=response) or 0.0 - except Exception: # noqa: BLE001 # unmapped judge model: the verdict still counts, cost stays 0 + except Exception: # noqa: BLE001 # unmapped model: the attempt still counts, cost stays 0 return 0.0 @@ -374,6 +386,32 @@ def _judge_user_prompt(conversation: str, response_a: str, response_b: str) -> s ) +def _job_spend_counter_key(job_id: str) -> str: + return f"spend:shadow_eval:{job_id}" + + +async def _job_spend_from_counter(counter_key: str, fallback_spend: float, max_budget: float) -> float: + """The leg's spend through the cross-pod counter the key budget gates read. The owner + degrades internally to the fill-time DB floor and raises only under fail-closed + enforcement, which the caller honors by skipping the sample.""" + from litellm.proxy.proxy_server import get_current_spend + + return await get_current_spend(counter_key=counter_key, fallback_spend=fallback_spend, max_budget=max_budget) + + +async def _add_job_spend_to_counter(counter_key: str, cost: float) -> None: + """Advance the counter the moment a cost is known, so even a lost row closes the gate. + Known failure mode: a Redis outage freezes the counter (the owner invalidates it), the + gate degrades to the fill floor, and overshoot grows to in-flight plus one TTL of + samples, the same degradation the key budget counters accept.""" + try: + from litellm.proxy.proxy_server import increment_spend_counter + + await increment_spend_counter(counter_key=counter_key, increment=cost) + except Exception as e: # noqa: BLE001 # attempt recording must proceed; the row stays truth and the fill floor gates + verbose_logger.warning("shadow_eval: spend counter increment failed for %s: %s", counter_key, e) + + async def _key_or_team_is_over_budget(metadata: Mapping[str, object]) -> bool: """Whether the shadowed key or its team is over budget, decided by the same owners the request path uses, so counter keys and thresholds can never drift from auth's. @@ -438,8 +476,8 @@ def _request_was_routed_by(request_metadata: Mapping[str, object], router_name: @dataclass(frozen=True, slots=True) class _CallFailure: - """A shadow or judge call that produced no usable response. cost carries any judge - spend the failed attempt still billed, so job-level judge_spend never undercounts.""" + """A shadow or judge call that produced no usable response. cost carries any spend + the failed call still billed, so job-level spend figures never undercount.""" error: str cost: float = 0.0 @@ -452,6 +490,7 @@ class _ShadowResponse: text: str model: str tier: str | None + cost: float @dataclass(frozen=True, slots=True) @@ -478,8 +517,10 @@ class ActiveShadowEvalJob(BaseModel): shadow_percentage: float judge_model: str max_turns: int + max_budget: float | None = None ends_at: datetime attempts: int = 0 + spend: float = 0.0 @field_validator("ends_at") @classmethod @@ -500,7 +541,7 @@ class ActiveShadowEvalJob(BaseModel): return self.baseline_model or self.router_name -def _as_active_job(record: object, attempts: int) -> ActiveShadowEvalJob | None: +def _as_active_job(record: object, attempts: int, spend: float) -> ActiveShadowEvalJob | None: """The sampling path's view of one job row, or None for a row it cannot sample: an unknown direction, or a reverse job with no baseline model to duplicate against. Failing closed here is what keeps the dispatch path total.""" @@ -509,7 +550,7 @@ def _as_active_job(record: object, attempts: int) -> ActiveShadowEvalJob | None: except ValidationError as e: verbose_logger.debug("shadow_eval: skipping unsamplable job row: %s", e) return None - return job.model_copy(update={"attempts": attempts}) + return job.model_copy(update={"attempts": attempts, "spend": spend}) # mutable-ok: pydantic update payload _jobs_cache: Final = InMemoryCache(max_size_in_memory=4, default_ttl=_JOBS_CACHE_TTL_SECONDS) @@ -524,12 +565,17 @@ class ShadowEvalLogger(CustomLogger): router_provider: Callable[[], "Router | None"] | None = None, prisma_provider: Callable[[], "PrismaClient | None"] | None = None, jobs_cache: InMemoryCache | None = None, + job_spend_reader: Callable[[str, float, float], Awaitable[float]] | None = None, + job_spend_writer: Callable[[str, float], Awaitable[None]] | None = None, ) -> None: """Providers are callables so the proxy's lazily-initialized globals are resolved - at call time, not at logger construction.""" + at call time, not at logger construction. The spend reader and writer wrap the + proxy's cross-pod spend counter; tests inject a plain in-memory pair.""" self._router_provider = router_provider or default_router_provider self._prisma_provider = prisma_provider or _default_prisma_provider self._jobs_cache = jobs_cache or _jobs_cache + self._read_job_spend = job_spend_reader or _job_spend_from_counter + self._write_job_spend = job_spend_writer or _add_job_spend_to_counter self._inflight_shadow_tasks: int = 0 # Starts per job since the last cache fill, never decremented within a # generation; the refill absorbs written rows and resets. @@ -556,18 +602,26 @@ class ShadowEvalLogger(CustomLogger): await prisma.db.litellm_shadowevalattempt.group_by( by=["job_id"], count=True, + sum={"judge_cost": True, "shadow_cost": True}, # mutable-ok: Prisma aggregate spec where={"job_id": {"in": [str(record.id) for record in records]}}, # mutable-ok: Prisma filter ) if records else () ) - attempt_counts: Final = {str(row["job_id"]): int(row["_count"]["_all"]) for row in grouped or []} + attempt_stats: Final = { # mutable-ok: frozen snapshot of the grouped read + str(row["job_id"]): ( + int(row["_count"]["_all"]), + float((row["_sum"] or {}).get("judge_cost") or 0.0) + + float((row["_sum"] or {}).get("shadow_cost") or 0.0), + ) + for row in grouped or [] + } by_key: Final = tuple( sorted( ( (str(record.api_key_id), job) for record in records or [] - if (job := _as_active_job(record, attempt_counts.get(str(record.id), 0))) is not None + if (job := _as_active_job(record, *attempt_stats.get(str(record.id), (0, 0.0)))) is not None ), key=itemgetter(0), ) @@ -624,6 +678,7 @@ class ShadowEvalLogger(CustomLogger): for job in (await self._active_jobs()).get(str(api_key_hash), ()) if datetime.now(timezone.utc) < job.ends_at and job.attempts + self._job_starts.get(job.id, 0) < job.max_turns + and (job.max_budget is None or job.spend < job.max_budget) and _sample_hits(request_id, job.id, job.shadow_percentage) and _request_was_routed_by(request_metadata, job.router_name) == (job.direction == "reverse") ) @@ -684,12 +739,28 @@ class ShadowEvalLogger(CustomLogger): return if await _key_or_team_is_over_budget(parent_metadata): return - + if job.max_budget is not None: + try: + spend: Final = await self._read_job_spend(_job_spend_counter_key(job.id), job.spend, job.max_budget) + except Exception as e: # noqa: BLE001 # unverifiable budget: skip the sample rather than spend on it + verbose_logger.warning("shadow_eval: budget unverifiable for %s, sample skipped: %s", job.id, e) + return + if spend >= job.max_budget: + return shadow: Final = await self._call_router_shadow(job.shadow_target, messages, shadow_params, parent_metadata) - if isinstance(shadow, _CallFailure): - await self._record_attempt(prisma, job, request_id, control_tier, outcome="error", error=shadow.error) - return - + except Exception as e: # noqa: BLE001 # detached task: nothing billed yet, record and never raise + verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e) + await self._record_attempt( + prisma, job, request_id, control_tier, outcome="error", error=f"pipeline error: {e}" + ) + return + if isinstance(shadow, _CallFailure): + await self._record_attempt( + prisma, job, request_id, control_tier, outcome="error", error=shadow.error, shadow_cost=shadow.cost + ) + return + # From here the shadow call has billed, so every exit records its cost. + try: verdict: Final = await self._call_judge( judge_model=job.judge_model, messages=messages, @@ -707,6 +778,7 @@ class ShadowEvalLogger(CustomLogger): error=verdict.error, shadow=shadow, judge_cost=verdict.cost, + shadow_cost=shadow.cost, ) return await self._record_attempt( @@ -719,15 +791,23 @@ class ShadowEvalLogger(CustomLogger): real_model=real_model, confidence=verdict.confidence, judge_cost=verdict.cost, + shadow_cost=shadow.cost, ) - except Exception as e: # noqa: BLE001 # detached task: record what happened, never raise + except Exception as e: # noqa: BLE001 # detached task: the shadow call billed, record its cost, never raise verbose_logger.debug("shadow_eval: pipeline failed for %s: %s", request_id, e) await self._record_attempt( - prisma, job, request_id, control_tier, outcome="error", error=f"pipeline error: {e}" + prisma, + job, + request_id, + control_tier, + outcome="error", + error=f"pipeline error: {e}", + shadow=shadow, + shadow_cost=shadow.cost, ) - @staticmethod async def _record_attempt( + self, prisma: "PrismaClient | None", job: ActiveShadowEvalJob, request_id: str, @@ -738,8 +818,11 @@ class ShadowEvalLogger(CustomLogger): real_model: str = "", confidence: float | None = None, judge_cost: float = 0.0, + shadow_cost: float = 0.0, error: str | None = None, ) -> None: + if judge_cost + shadow_cost > 0: + await self._write_job_spend(_job_spend_counter_key(job.id), judge_cost + shadow_cost) if prisma is None: return try: @@ -753,6 +836,7 @@ class ShadowEvalLogger(CustomLogger): "shadow_model": shadow.model if shadow else None, "confidence": confidence, "judge_cost": judge_cost, + "shadow_cost": shadow_cost, "error": error[:_MAX_ERROR_CHARS] if error else None, } ) @@ -792,11 +876,12 @@ class ShadowEvalLogger(CustomLogger): return _CallFailure(f"shadow router call failed: {_failure_detail(e)}") text: Final = _chat_final_text(response) if not text: - return _CallFailure("shadow router returned an empty response") + return _CallFailure("shadow router returned an empty response", cost=_call_cost(response)) return _ShadowResponse( text=text, model=str(getattr(response, "model", None) or _routing_decision(shadow_metadata).get("routed_model") or ""), tier=_routed_tier(shadow_metadata), + cost=_call_cost(response), ) async def _call_judge( @@ -843,11 +928,11 @@ class ShadowEvalLogger(CustomLogger): verdict: Final = PairwiseVerdict.model_validate(parse_json_verdict(raw)) except Exception as e: # noqa: BLE001 # malformed verdicts become error rows verbose_logger.debug("shadow_eval: unparseable judge verdict: %s", e) - return _CallFailure(f"unparseable judge verdict: {e}", cost=_judge_call_cost(response)) + return _CallFailure(f"unparseable judge verdict: {e}", cost=_call_cost(response)) return _JudgeVerdict( preference=_unmask_preference(verdict.preference, real_is_a), confidence=max(0.0, min(1.0, verdict.confidence)), - cost=_judge_call_cost(response), + cost=_call_cost(response), ) diff --git a/litellm/litellm_core_utils/env_utils.py b/litellm/litellm_core_utils/env_utils.py index af0520eaf31..d641884b4cd 100644 --- a/litellm/litellm_core_utils/env_utils.py +++ b/litellm/litellm_core_utils/env_utils.py @@ -2,6 +2,7 @@ Utility helpers for reading and parsing environment variables. """ +import logging import os from typing import Final @@ -22,6 +23,26 @@ def get_env_int(env_var: str, default: int) -> int: return default +def get_env_int_in_range(env_var: str, default: int, minimum: int, maximum: int) -> int: + """Parse an environment variable as an integer constrained to ``[minimum, maximum]``. + + Values outside the range fall back to the default and warn, so a misconfigured knob can + neither crash the caller nor silently change the meaning of what it computes. + """ + value: Final = get_env_int(env_var, default) + if minimum <= value <= maximum: + return value + logging.getLogger("LiteLLM").warning( + "%s=%s is outside the supported range [%s, %s]. Falling back to %s.", + env_var, + value, + minimum, + maximum, + default, + ) + return default + + def get_env_int_or_none(env_var: str) -> int | None: """Parse an environment variable as an integer, returning None when it is unset or unusable. diff --git a/litellm/litellm_core_utils/get_litellm_params.py b/litellm/litellm_core_utils/get_litellm_params.py index 3eb8c163d5c..b12c715c9f5 100644 --- a/litellm/litellm_core_utils/get_litellm_params.py +++ b/litellm/litellm_core_utils/get_litellm_params.py @@ -21,6 +21,14 @@ AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset( } ) +# The per-deployment Rust opt-in. +RUST_KWARG_KEY: Final = "rust" + +# Keys `completion()` forwards from its own kwargs into `get_litellm_params`, +# which are otherwise invisible to it because that call site passes explicit +# named arguments rather than `**kwargs`. +FORWARDED_KWARGS_KEYS: Final = AWS_CREDENTIAL_KWARGS_KEYS | frozenset({RUST_KWARG_KEY}) + # Pre-define optional kwargs keys as frozenset for O(1) lookups # These are extracted from kwargs only if present, avoiding unnecessary .get() calls OPTIONAL_KWARGS_KEYS: Final = ( @@ -47,6 +55,10 @@ OPTIONAL_KWARGS_KEYS: Final = ( "itpm", "otpm", "use_xai_oauth", + # The per-deployment Rust opt-in. `all_litellm_params` keeps it out + # of the provider body; this keeps it *in* litellm_params, which is + # where the chat completions handlers read it from. + RUST_KWARG_KEY, } ) | AWS_CREDENTIAL_KWARGS_KEYS diff --git a/litellm/litellm_core_utils/ptu_pricing.py b/litellm/litellm_core_utils/ptu_pricing.py index a1f8bb36e27..6923e6beb96 100644 --- a/litellm/litellm_core_utils/ptu_pricing.py +++ b/litellm/litellm_core_utils/ptu_pricing.py @@ -79,6 +79,45 @@ def _as_utc(value: object) -> datetime | None: return None +def _named(reason: str, model_name: str | None) -> str: + """The reason on its own for a caller that already has the deployment in hand, else named.""" + return reason if model_name is None else f"PTU configuration on model '{model_name}' is invalid: {reason}" + + +def ptu_config_error(model_info: Mapping[str, object], *, model_name: str | None = None) -> str | None: + """Why this PTU configuration cannot be honoured, else None. + + Both the model endpoints and config.yaml registration ask this, so a deployment that + one refuses is refused by the other for the same stated reason. + + Window ordering is checked before the count/rate gate. A patch that touches only one end + of the window carries no count or rate, so leaving the order to that gate would let an + inverted window reach the row; the next load then fails to parse it and drops the + deployment out of the router, where no further patch can repair it. + """ + effective_from: Final = _as_utc(model_info.get("ptu_effective_from")) + effective_to: Final = _as_utc(model_info.get("ptu_effective_to")) + if effective_from is not None and effective_to is not None and effective_to <= effective_from: + return _named("ptu_effective_to must be after ptu_effective_from", model_name) + + has_count: Final = model_info.get("ptu_count") is not None + has_rate: Final = model_info.get("cost_per_ptu_per_hour") is not None + if not has_count and not has_rate: + return None + if has_count != has_rate: + return _named("ptu_count and cost_per_ptu_per_hour must be set together", model_name) + if effective_from is None: + return _named( + "ptu_effective_from is required when PTU fields are set. Flat cost accrues from that " + "instant, so without it the start would have to be inferred and a deployment configured " + "today could be billed for days it did not exist", + model_name, + ) + if not model_info.get("team_id"): + return _named("team_id is required when PTU fields are set (one model maps to one team)", model_name) + return None + + def ptu_terms(model_info: Mapping[str, object]) -> PTUTerms | None: """The reservation this deployment accrues flat cost for, else None. diff --git a/litellm/litellm_core_utils/thread_pool_executor.py b/litellm/litellm_core_utils/thread_pool_executor.py index 881a91400df..f989f20247f 100644 --- a/litellm/litellm_core_utils/thread_pool_executor.py +++ b/litellm/litellm_core_utils/thread_pool_executor.py @@ -1,6 +1,82 @@ -from concurrent.futures import ThreadPoolExecutor -from typing import Final +import logging +import threading +import time +from collections.abc import Callable +from concurrent.futures import Future, ThreadPoolExecutor +from typing import Final, ParamSpec, TypeVar -MAX_THREADS: Final = 100 -# Create a ThreadPoolExecutor -executor: Final = ThreadPoolExecutor(max_workers=MAX_THREADS) +from litellm._logging import verbose_logger +from litellm.constants import ( + LOGGING_EXECUTOR_DROPPED_TASK_LOG_INTERVAL_SECONDS, + LOGGING_EXECUTOR_MAX_PENDING_TASKS, + LOGGING_EXECUTOR_MAX_THREADS, +) + +MAX_THREADS: Final = LOGGING_EXECUTOR_MAX_THREADS + +_P = ParamSpec("_P") +_T = TypeVar("_T") + + +class BoundedLoggingThreadPoolExecutor(ThreadPoolExecutor): + """ThreadPoolExecutor with a cap on queued-plus-running tasks. + + The default ThreadPoolExecutor work queue is unbounded, and every queued + logging task pins its request/response payload in memory, so a sustained + burst of sync callbacks slower than request arrival grows memory without + bound. Logging is best-effort: once the cap is reached, new submissions + are dropped with a rate-limited warning instead of queueing forever. + """ + + def __init__( + self, + max_workers: int, + max_pending_tasks: int, + drop_log_interval_seconds: float = LOGGING_EXECUTOR_DROPPED_TASK_LOG_INTERVAL_SECONDS, + logger: logging.Logger = verbose_logger, + ) -> None: + super().__init__(max_workers=max_workers, thread_name_prefix="litellm-logging") + self._max_pending_tasks: Final = max_pending_tasks + self._drop_log_interval_seconds: Final = drop_log_interval_seconds + self._logger: Final = logger + self._pending_slots: Final = threading.Semaphore(max_pending_tasks) + self._drop_lock: Final = threading.Lock() + self._dropped_since_last_log = 0 + self._last_drop_log_time = 0.0 + + def submit(self, fn: Callable[_P, _T], /, *args: _P.args, **kwargs: _P.kwargs) -> Future[_T]: + if not self._pending_slots.acquire(blocking=False): + self._record_drop() + dropped_future: Final[Future[_T]] = Future() + dropped_future.cancel() + return dropped_future + try: + future: Final = super().submit(fn, *args, **kwargs) + except BaseException: + self._pending_slots.release() + raise + future.add_done_callback(lambda _: self._pending_slots.release()) + return future + + def _record_drop(self) -> None: + with self._drop_lock: + self._dropped_since_last_log += 1 + now: Final = time.monotonic() + if now - self._last_drop_log_time < self._drop_log_interval_seconds: + return + dropped_count: Final = self._dropped_since_last_log + self._dropped_since_last_log = 0 + self._last_drop_log_time = now + + self._logger.warning( + "litellm logging executor backlog is full (max_pending_tasks=%s); dropped %s logging task(s) " + "since the last warning. Set LOGGING_EXECUTOR_MAX_PENDING_TASKS to raise the cap.", + self._max_pending_tasks, + dropped_count, + ) + + +executor: Final = BoundedLoggingThreadPoolExecutor( + max_workers=MAX_THREADS, + max_pending_tasks=LOGGING_EXECUTOR_MAX_PENDING_TASKS, +) diff --git a/litellm/litellm_core_utils/token_counter.py b/litellm/litellm_core_utils/token_counter.py index 17f3dea72ec..858b078d626 100644 --- a/litellm/litellm_core_utils/token_counter.py +++ b/litellm/litellm_core_utils/token_counter.py @@ -19,6 +19,7 @@ from litellm.constants import ( MAX_SHORT_SIDE_FOR_IMAGE_HIGH_RES, MAX_TILE_HEIGHT, MAX_TILE_WIDTH, + TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS, ) from litellm.litellm_core_utils.default_encoding import encoding as default_encoding from litellm.litellm_core_utils.url_utils import safe_get @@ -305,6 +306,16 @@ Type for a function that counts tokens in a string. """ +def _get_tiktoken_count_function( + encode_length: Callable[[str], int], + chunk_size: int = TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS, +) -> TokenCounterFunction: + def count_tokens(text: str) -> int: + return sum(encode_length(text[start : start + chunk_size]) for start in range(0, len(text), chunk_size)) + + return count_tokens + + class _MessageCountParams: """ A class to hold the parameters for counting tokens in messages. @@ -531,6 +542,7 @@ def _get_count_function( enc: Final = tokenizer_json["tokenizer"].encode(text) return len(enc.ids) + return count_tokens elif tokenizer_json["type"] == "openai_tokenizer": model_to_use: Final = _fix_model_name(model) try: @@ -542,17 +554,18 @@ def _get_count_function( print_verbose("Warning: model not found. Using cl100k_base encoding.") encoding = tiktoken.get_encoding("cl100k_base") - def count_tokens(text: str) -> int: + def encode_length(text: str) -> int: return len(encoding.encode(text, disallowed_special=())) + return _get_tiktoken_count_function(encode_length) else: raise ValueError("Unsupported tokenizer type") else: - def count_tokens(text: str) -> int: + def encode_length(text: str) -> int: return len(default_encoding.encode(text, disallowed_special=())) - return count_tokens + return _get_tiktoken_count_function(encode_length) def _fix_model_name(model: str) -> str: diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index 39d3947c07c..d9bb0d7abff 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -24,6 +24,8 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, ) +from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge +from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts from litellm.types.llms.anthropic import ( ContentBlockDelta, ContentBlockStart, @@ -361,30 +363,135 @@ class AnthropicChatCompletion(BaseLLM): if config is None: raise ValueError(f"Provider config not found for model: {model} and provider: {custom_llm_provider}") - data = config.transform_request( + def build_request() -> tuple[dict, dict]: # mutable-ok: rewritten in place downstream + """Translate the request the Python way, returning `(headers, data)`. + + The pair stays mutable because the streaming path rewrites it in + place (`data["stream"] = True`) before sending. + + Shared by the normal path and by the Rust path's fallback, which + builds it only when the Rust call did not serve the request. + """ + request_data: Final = config.transform_request( + model=model, + messages=messages, + optional_params={**optional_params, "is_vertex_request": is_vertex_request}, + litellm_params=litellm_params, + headers=headers, + ) + return update_request_with_filtered_beta( + headers=headers, + request_data=request_data, + provider=custom_llm_provider, + ) + + # The Rust core owns the whole call for the subset it accepts, so ask + # before transforming: whichever path runs emits pre_call exactly once. + # `get_config` merges the class-level defaults (Anthropic's required + # `max_tokens` among them) that `transform_request` would have applied. + rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy + **AnthropicConfig.get_config(model=model), + **optional_params, + } + serves_via_rust: Final = rust_chat_completions_accepts( model=model, messages=messages, - optional_params={**optional_params, "is_vertex_request": is_vertex_request}, + optional_params=rust_optional_params, + custom_llm_provider=custom_llm_provider, litellm_params=litellm_params, - headers=headers, + stream=stream, ) - - headers, data = update_request_with_filtered_beta( - headers=headers, - request_data=data, - provider=custom_llm_provider, - ) - - ## LOGGING - logging_obj.pre_call( - input=messages, - api_key=api_key, - additional_args={ - "complete_input_dict": data, + if serves_via_rust: + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "model": model, + "messages": messages, + **rust_optional_params, + }, "api_base": api_base, "headers": headers, - }, - ) + } + logging_obj.pre_call(input=messages, api_key=api_key, additional_args=rust_logging_args) + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key=api_key, + additional_args=rust_logging_args, + ) + if acompletion is True: + + async def python_fallback() -> "ModelResponse | CustomStreamWrapper": + # pre_call already fired for this request above. The Rust + # path only declines before the provider is called, so this + # is the same attempt continuing, not a second one. + fallback_headers, fallback_data = build_request() + return await self.acompletion_function( + model=model, + messages=messages, + data=fallback_data, + api_base=api_base, + custom_prompt_dict=custom_prompt_dict, + model_response=model_response, + print_verbose=print_verbose, + encoding=encoding, + api_key=api_key, + provider_config=config, + logging_obj=logging_obj, + optional_params=optional_params, + stream=stream, + _is_function_call=_is_function_call, + litellm_params=litellm_params, + logger_fn=logger_fn, + headers=fallback_headers, + client=client, + json_mode=json_mode, + timeout=timeout, + ) + + return rust_chat_completions_bridge.achat_completions_or_fallback( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + python_fallback=python_fallback, + ) + rust_response: Final = rust_chat_completions_bridge.chat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + ) + if rust_response is not None: + return rust_response + + headers, data = build_request() + + ## LOGGING + # Reaching here with `serves_via_rust` set means the Rust attempt + # declined at call time, before the provider was called, and already + # logged this request. That is the same attempt continuing. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, + api_key=api_key, + additional_args={ + "complete_input_dict": data, + "api_base": api_base, + "headers": headers, + }, + ) print_verbose(f"_is_function_call: {_is_function_call}") if acompletion is True: if ( diff --git a/litellm/llms/bedrock/chat/converse_handler.py b/litellm/llms/bedrock/chat/converse_handler.py index 25e544f4521..3a093e1f939 100644 --- a/litellm/llms/bedrock/chat/converse_handler.py +++ b/litellm/llms/bedrock/chat/converse_handler.py @@ -14,6 +14,8 @@ from litellm.llms.custom_httpx.http_handler import ( _get_httpx_client, get_async_httpx_client, ) +from litellm.rust_bridge import chat_completions as rust_chat_completions_bridge +from litellm.rust_bridge.chat_completions import rust_chat_completions_accepts from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -169,6 +171,7 @@ class BedrockConverseLLM(BaseAWSLLM): headers: dict = {}, client: AsyncHTTPHandler | None = None, api_key: str | None = None, + skip_pre_call_logging: bool = False, ) -> ModelResponse | CustomStreamWrapper: request_data: Final = await litellm.AmazonConverseConfig()._async_transform_request( model=model, @@ -190,15 +193,19 @@ class BedrockConverseLLM(BaseAWSLLM): ) ## LOGGING - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": api_base, - "headers": prepped.headers, - }, - ) + # The Rust path already logged this request's pre_call before handing + # it here, and it only declines before the provider is called, so this + # is the same attempt continuing rather than a second one. + if not skip_pre_call_logging: + logging_obj.pre_call( + input=messages, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": api_base, + "headers": prepped.headers, + }, + ) headers = dict(prepped.headers) if client is None or not isinstance(client, AsyncHTTPHandler): @@ -354,6 +361,94 @@ class BedrockConverseLLM(BaseAWSLLM): # Filter beta headers in HTTP headers before making the request headers = update_headers_with_filtered_beta(headers=headers, provider="bedrock_converse") + + # The Rust core owns the whole call for the subset it accepts. Ask + # before transforming so whichever path runs emits pre_call once, and + # hand down the credentials, region and endpoint this handler already + # resolved so both paths sign as the same principal. + rust_optional_params: Final = { # mutable-ok: json.dumps in the bridge rejects a mappingproxy + **optional_params, + **{ # mutable-ok: merged into its mutable parent above + key: value + for key, value in ( + ("aws_access_key_id", credentials.access_key), + ("aws_secret_access_key", credentials.secret_key), + ("aws_session_token", credentials.token), + ("aws_region_name", aws_region_name), + ) + if value is not None + }, + } + serves_via_rust: Final = rust_chat_completions_accepts( + model=model, + messages=messages, + optional_params=rust_optional_params, + custom_llm_provider="bedrock", + litellm_params=litellm_params, + stream=stream, + ) + if serves_via_rust: + rust_logging_args: Final = { # mutable-ok: logging callbacks read additional_args as a plain dict + "complete_input_dict": { # mutable-ok: same, and it is serialized alongside its parent + "messages": messages, + **optional_params, + }, + "api_base": proxy_endpoint_url, + "headers": headers, + } + logging_obj.pre_call(input=messages, api_key="", additional_args=rust_logging_args) + log_rust_post_call: Final = rust_chat_completions_bridge.response_logger( + logging_obj=logging_obj, + messages=messages, + api_key="", + additional_args=rust_logging_args, + ) + if acompletion: + return rust_chat_completions_bridge.achat_completions_or_fallback( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + python_fallback=lambda: self.async_completion( + model=model, + messages=messages, + api_base=proxy_endpoint_url, + model_response=model_response, + encoding=encoding, + logging_obj=logging_obj, + optional_params=optional_params, + stream=stream, + litellm_params=litellm_params, + logger_fn=logger_fn, + headers=headers, + timeout=timeout, + client=client, + credentials=credentials, + api_key=api_key, + skip_pre_call_logging=True, + ), + ) + rust_response: Final = rust_chat_completions_bridge.chat_completions( + model=model, + messages=messages, + optional_params=rust_optional_params, + model_response=model_response, + api_key=api_key, + api_base=proxy_endpoint_url, + custom_llm_provider="bedrock", + extra_headers=headers, + timeout=timeout, + on_response=log_rust_post_call, + ) + if rust_response is not None: + return rust_response + ### ROUTING (ASYNC, STREAMING, SYNC) if acompletion: if isinstance(client, HTTPHandler): @@ -420,15 +515,21 @@ class BedrockConverseLLM(BaseAWSLLM): ) ## LOGGING - logging_obj.pre_call( - input=messages, - api_key="", - additional_args={ - "complete_input_dict": data, - "api_base": proxy_endpoint_url, - "headers": prepped.headers, - }, - ) + # Reaching here with `serves_via_rust` set means the synchronous Rust + # attempt declined at call time, before the provider was called, and + # already logged this request. That is the same attempt continuing. + # The asynchronous branch above returns before this point, and hands + # its own fallback `skip_pre_call_logging=True` for the same reason. + if not serves_via_rust: + logging_obj.pre_call( + input=messages, + api_key="", + additional_args={ + "complete_input_dict": data, + "api_base": proxy_endpoint_url, + "headers": prepped.headers, + }, + ) if client is None or isinstance(client, AsyncHTTPHandler): _params: Final = {} if timeout is not None: diff --git a/litellm/llms/custom_httpx/container_handler.py b/litellm/llms/custom_httpx/container_handler.py index 7690351e3b2..91d68aa3bfb 100644 --- a/litellm/llms/custom_httpx/container_handler.py +++ b/litellm/llms/custom_httpx/container_handler.py @@ -39,6 +39,10 @@ RESPONSE_TYPES: Final[dict[str, type]] = { "DeleteContainerFileResponse": DeleteContainerFileResponse, } +ContainerEndpointResponse = ( + ContainerFileListResponse | ContainerFileObject | DeleteContainerFileResponse | bytes | dict[str, object] +) + def _load_endpoints_config() -> dict: """Load the endpoints configuration from JSON file.""" @@ -101,6 +105,51 @@ def _build_query_params( return params +def _error_message_from_response(response: httpx.Response) -> str: + try: + body: Final = response.json() + except ValueError: + return response.text + + if isinstance(body, dict) and isinstance(body.get("error"), dict): + message: Final = body["error"].get("message") + if isinstance(message, str): + return message + + return response.text + + +def _transform_response( + response: httpx.Response, + returns_binary: bool, + response_type_name: str, +) -> ContainerEndpointResponse: + from litellm.llms.base_llm.chat.transformation import BaseLLMException + + if httpx.codes.is_error(response.status_code): + raise BaseLLMException( + status_code=response.status_code, + message=_error_message_from_response(response), + headers=dict(response.headers), + ) + + if returns_binary: + return response.content + + response_json: Final = response.json() + if "error" in response_json: + raise BaseLLMException( + status_code=response.status_code, + message=response_json.get("error", {}).get("message", str(response_json)), + headers=dict(response.headers), + ) + + response_type: Final = RESPONSE_TYPES.get(response_type_name) + if response_type: + return response_type(**response_json) + return response_json + + def _prepare_multipart_file_upload( file: Any, headers: dict[str, Any], @@ -270,27 +319,11 @@ class GenericContainerHandler: else: raise ValueError(f"Unsupported HTTP method: {method}") - # For binary responses, return raw content - if returns_binary: - return response.content - - # Check for error response - response_json: Final = response.json() - if "error" in response_json: - from litellm.llms.base_llm.chat.transformation import BaseLLMException - - error_msg: Final = response_json.get("error", {}).get("message", str(response_json)) - raise BaseLLMException( - status_code=response.status_code, - message=error_msg, - headers=dict(response.headers), - ) - - # Parse response - response_type: Final = RESPONSE_TYPES.get(endpoint_config["response_type"]) - if response_type: - return response_type(**response_json) - return response_json + return _transform_response( + response=response, + returns_binary=returns_binary, + response_type_name=endpoint_config["response_type"], + ) except Exception as e: raise e @@ -378,27 +411,11 @@ class GenericContainerHandler: else: raise ValueError(f"Unsupported HTTP method: {method}") - # For binary responses, return raw content - if returns_binary: - return response.content - - # Check for error response - response_json: Final = response.json() - if "error" in response_json: - from litellm.llms.base_llm.chat.transformation import BaseLLMException - - error_msg: Final = response_json.get("error", {}).get("message", str(response_json)) - raise BaseLLMException( - status_code=response.status_code, - message=error_msg, - headers=dict(response.headers), - ) - - # Parse response - response_type: Final = RESPONSE_TYPES.get(endpoint_config["response_type"]) - if response_type: - return response_type(**response_json) - return response_json + return _transform_response( + response=response, + returns_binary=returns_binary, + response_type_name=endpoint_config["response_type"], + ) except Exception as e: raise e diff --git a/litellm/llms/perplexity/cost_calculator.py b/litellm/llms/perplexity/cost_calculator.py index 337fa8e630d..27835ecbfe8 100644 --- a/litellm/llms/perplexity/cost_calculator.py +++ b/litellm/llms/perplexity/cost_calculator.py @@ -21,14 +21,19 @@ def cost_per_token(model: str, usage: Usage) -> tuple[float, float]: Tuple[float, float] - prompt_cost_in_usd, completion_cost_in_usd """ ## USE PRE-CALCULATED COST FROM PERPLEXITY IF AVAILABLE - ## Perplexity returns accurate cost in usage.cost.total_cost including request fees + ## Perplexity returns accurate cost in usage.cost.total_cost including request fees. + ## By the time it reaches here, ResponseAPIUsage.parse_cost has already flattened + ## that dict down to a float, so both shapes must be accepted. cost_info: Final = getattr(usage, "cost", None) - if cost_info is not None and isinstance(cost_info, dict): - total_cost: Final = cost_info.get("total_cost") - if total_cost is not None: - # Return total cost as completion_cost (prompt_cost=0) since Perplexity - # doesn't break down by input/output in their cost object - return (0.0, float(total_cost)) + total_cost: float | None = None + if isinstance(cost_info, dict): + total_cost = cost_info.get("total_cost") + elif isinstance(cost_info, (int, float)) and not isinstance(cost_info, bool): + total_cost = float(cost_info) + if total_cost is not None: + # Return total cost as completion_cost (prompt_cost=0) since Perplexity + # doesn't break down by input/output in their cost object + return (0.0, float(total_cost)) ## FALLBACK: Calculate cost manually if Perplexity doesn't provide it ## GET MODEL INFO diff --git a/litellm/main.py b/litellm/main.py index c3af24e1a51..7cfd322f3d0 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -75,7 +75,7 @@ from litellm.litellm_core_utils.chat_completion_agentic_loop import ( from litellm.litellm_core_utils.completion_timeout import CompletionTimeout from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_litellm_params import ( - AWS_CREDENTIAL_KWARGS_KEYS, + FORWARDED_KWARGS_KEYS, OPTIONAL_KWARGS_KEYS, ) from litellm.litellm_core_utils.get_provider_specific_headers import ( @@ -5451,7 +5451,7 @@ def completion( tpm=kwargs.get("tpm"), rpm=kwargs.get("rpm"), use_xai_oauth=kwargs.get("use_xai_oauth", False), - **{key: kwargs[key] for key in AWS_CREDENTIAL_KWARGS_KEYS if key in kwargs}, + **{key: kwargs[key] for key in FORWARDED_KWARGS_KEYS if key in kwargs}, ) cast(LiteLLMLoggingObj, logging).update_environment_variables( model=model, diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b9c8824aa67..858eab672e5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -34758,6 +34758,50 @@ "supports_reasoning": false, "supports_function_calling": true }, + "perplexity/perplexity/deepseek-v4-flash-0731": { + "cache_read_input_token_cost": 2.8e-08, + "input_cost_per_token": 1.3e-07, + "litellm_provider": "perplexity", + "mode": "responses", + "output_cost_per_token": 2.6e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models", + "supports_web_search": true, + "supports_reasoning": true, + "supports_function_calling": true + }, + "perplexity/perplexity/glm-5.2": { + "cache_read_input_token_cost": 1.4e-07, + "input_cost_per_token": 1.4e-06, + "litellm_provider": "perplexity", + "mode": "responses", + "output_cost_per_token": 4.4e-06, + "source": "https://docs.perplexity.ai/docs/agent-api/models", + "supports_web_search": true, + "supports_reasoning": true, + "supports_function_calling": true + }, + "perplexity/perplexity/kimi-k3": { + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "perplexity", + "mode": "responses", + "output_cost_per_token": 1.5e-05, + "source": "https://docs.perplexity.ai/docs/agent-api/models", + "supports_web_search": true, + "supports_reasoning": true, + "supports_function_calling": true + }, + "perplexity/perplexity/kimi-k2.7-code": { + "cache_read_input_token_cost": 1.9e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "perplexity", + "mode": "responses", + "output_cost_per_token": 4e-06, + "source": "https://docs.perplexity.ai/docs/agent-api/models", + "supports_web_search": true, + "supports_reasoning": false, + "supports_function_calling": true + }, "perplexity/pplx-embed-v1-0.6b": { "input_cost_per_token": 4e-09, "litellm_provider": "perplexity", diff --git a/litellm/proxy/_experimental/mcp_server/db.py b/litellm/proxy/_experimental/mcp_server/db.py index 2f07a8b716c..28638ed9c77 100644 --- a/litellm/proxy/_experimental/mcp_server/db.py +++ b/litellm/proxy/_experimental/mcp_server/db.py @@ -1224,14 +1224,28 @@ def _decode_user_credential(stored: str) -> str | None: return None -def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None: - """Return the OAuth2 payload dict if ``stored`` holds one, else ``None``. +def _warn_undecryptable_credential(user_id: str, server_id: str) -> None: + """Log the one credential state that otherwise reads as "user never authorized".""" + verbose_proxy_logger.warning( + "MCP user credential for user=%s server=%s could not be decrypted (likely written under a " + "previous LITELLM_SALT_KEY); the user is treated as not connected and must re-authorize.", + user_id, + server_id, + ) + + +def _parse_oauth_payload(decoded: str | None) -> OAuthCredentialPayload | None: + """Return the OAuth2 payload dict if ``decoded`` holds one, else ``None``. A row is considered an OAuth2 credential iff its decoded value parses as a JSON object with ``"type": "oauth2"``. Plain BYOK credentials (which share the same column) decode to a non-JSON string and return ``None``. + + Callers that need to tell an unreadable row from a readable non-OAuth2 one + pass the result of :func:`_decode_user_credential` so a single decode + answers both questions: ``None`` there means the value can be neither + decrypted nor base64-decoded, so no caller can ever recover it. """ - decoded: Final = _decode_user_credential(stored) if decoded is None: return None parsed: OAuthCredentialPayload | None @@ -1244,6 +1258,11 @@ def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None: return None +def _decode_oauth_payload(stored: str) -> OAuthCredentialPayload | None: + """Return the OAuth2 payload dict held in ``stored``, else ``None``.""" + return _parse_oauth_payload(_decode_user_credential(stored)) + + async def rotate_mcp_user_credentials_master_key(prisma_client: PrismaClient, new_master_key: str): """Re-encrypt every ``LiteLLM_MCPUserCredentials`` row with ``new_master_key``. @@ -1415,15 +1434,25 @@ async def store_user_oauth_credential( # (e.g. during token refresh), saving an extra DB round-trip. if not skip_byok_guard: existing: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id) - if existing is not None and _decode_oauth_payload(existing.credential_b64) is None: - # Existing row is either a BYOK secret or an OAuth2 row that no - # longer decrypts (e.g. after a salt-key rotation). In either - # case, refuse to overwrite — the caller would clobber data - # that may still be recoverable. - raise ValueError( - f"Existing credential for user {user_id} and server " - f"{server_id} could not be verified as an OAuth2 token. " - f"Refusing to overwrite." + decoded: Final = _decode_user_credential(existing.credential_b64) if existing is not None else None + if existing is not None and _parse_oauth_payload(decoded) is None: + # Refuse only while the row still holds readable content, which is a live BYOK + # secret that overwriting would destroy. A row that does not decode was written + # under a different LITELLM_SALT_KEY, and one that decodes to nothing holds no + # secret at all; refusing either preserves nothing and instead wedges the user + # out of the OAuth flow for good, since re-authorizing is their only recovery. + if decoded: + raise ValueError( + f"Existing credential for user {user_id} and server " + f"{server_id} could not be verified as an OAuth2 token. " + f"Refusing to overwrite." + ) + verbose_proxy_logger.warning( + "store_user_oauth_credential: existing credential for user=%s server=%s could not be " + "decrypted (likely written under a previous LITELLM_SALT_KEY); replacing it with the " + "newly authorized OAuth2 token.", + user_id, + server_id, ) encoded: Final = encrypt_value_helper(json.dumps(payload)) @@ -1461,7 +1490,10 @@ async def get_user_oauth_credential( row: Final = await _db_find_user_credential_row(prisma_client, user_id, server_id) if row is None: return None - return _decode_oauth_payload(row.credential_b64) + decoded: Final = _decode_user_credential(row.credential_b64) + if decoded is None: + _warn_undecryptable_credential(user_id, server_id) + return _parse_oauth_payload(decoded) async def list_user_oauth_credentials( @@ -1473,7 +1505,10 @@ async def list_user_oauth_credentials( rows: Final = await _db_find_user_credential_rows(prisma_client, {"user_id": user_id}) results: Final[list[OAuthCredentialPayload]] = [] for row in rows: - payload = _decode_oauth_payload(row.credential_b64) + decoded = _decode_user_credential(row.credential_b64) + if decoded is None: + _warn_undecryptable_credential(user_id, row.server_id) + payload = _parse_oauth_payload(decoded) if payload is None: continue payload["server_id"] = row.server_id diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 26a6f8d1251..dbe97dd5bce 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -46,7 +46,7 @@ from litellm.constants import ( MCP_TOOL_LISTING_TIMEOUT, ) from litellm.exceptions import BlockedPiiEntityError, GuardrailRaisedException -from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth +from litellm.experimental_mcp_client.client import MCPClient, MCPSigV4Auth, strip_auth_scheme from litellm.integrations.custom_guardrail import ( _sync_guardrail_info_to_logging_obj, # pyright: ignore[reportPrivateUsage] - the same bridge @log_guardrail_information uses; reimplementing it here would fork the metadata-key logic ) @@ -841,12 +841,17 @@ def _without_authorization( def _format_byok_openapi_auth_header(mcp_server: MCPServer, mcp_auth_header: str) -> str: - """Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection.""" + """Format a raw BYOK credential for OpenAPI tool ``Authorization`` injection. + + A non-BYOK server short-circuits ``_resolve_byok_mcp_auth_header``, so the value here can also + be the deprecated global ``x-mcp-auth``, which is a complete header value and would otherwise + be given a second scheme. + """ if mcp_server.auth_type == MCPAuth.api_key: - return f"ApiKey {mcp_auth_header}" + return f"ApiKey {strip_auth_scheme(mcp_auth_header, 'ApiKey')}" if mcp_server.auth_type == MCPAuth.basic: - return f"Basic {mcp_auth_header}" - return f"Bearer {mcp_auth_header}" + return f"Basic {strip_auth_scheme(mcp_auth_header, 'Basic')}" + return f"Bearer {strip_auth_scheme(mcp_auth_header, 'Bearer')}" def _openapi_forwarded_extra_headers( diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 05fc6e07176..669924077e3 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -515,15 +515,28 @@ class LiteLLMRoutes(enum.Enum): # allowed_routes=["mcp_routes"], which should cover both halves. mcp_routes = mcp_inference_routes + mcp_management_routes - agent_routes = [ - "/v1/agents", - "/v1/agents/{agent_id}", + # A2A agent invocation / discovery routes — data-plane. Gated by DISABLE_LLM_API_ENDPOINTS. + agent_inference_routes = ( "/agents", "/a2a/{agent_id}", "/a2a/{agent_id}/message/send", "/a2a/{agent_id}/message/stream", "/a2a/{agent_id}/.well-known/agent-card.json", - ] + ) + + # Agent registry CRUD routes — control-plane. Gated by DISABLE_ADMIN_ENDPOINTS. + # The handlers in agent_endpoints/endpoints.py enforce proxy-admin on writes and + # scope reads by role, so these also appear in self_managed_routes. + agent_management_routes = ( + "/v1/agents", + "/v1/agents/{agent_id}", + "/v1/agents/make_public", + "/v1/agents/{agent_id}/make_public", + ) + + # Backwards-compat union — virtual keys may be configured with + # allowed_routes=["agent_routes"], which should cover both halves. + agent_routes = agent_inference_routes + agent_management_routes google_routes = [ "/v1beta/models/{model_name:path}:countTokens", @@ -563,7 +576,7 @@ class LiteLLMRoutes(enum.Enum): + apply_guardrail_routes + mcp_inference_routes + litellm_native_routes - + agent_routes + + list(agent_inference_routes) + model_info_routes ) info_routes = [ @@ -664,6 +677,7 @@ class LiteLLMRoutes(enum.Enum): ] + key_management_routes + mcp_management_routes + + list(agent_management_routes) ) spend_tracking_routes = [ @@ -836,6 +850,9 @@ class LiteLLMRoutes(enum.Enum): # proxy admin, or team admin naming their own team via team_id "/auto_router/test_routing", "/auto_router/validate_complexity_router_config", + # Agent registry - reads are role-scoped and writes are proxy-admin-gated + # inside agent_endpoints/endpoints.py + *agent_management_routes, ] # routes that manage their own allowed/disallowed logic ## Org Admin Routes ## @@ -2551,6 +2568,14 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): None, description="Maximum retention period for auto-router benchmark session rollup rows (e.g., '365d'). Rows whose last turn is older than this are deleted by the spend log cleanup job, on that job's schedule. Unset means rollup rows are never deleted.", ) + maximum_health_check_retention_period: str | None = Field( + None, + description=( + "Maximum retention period for health-check rows (e.g., '30d'). Rows whose checked_at is older than this " + "are deleted by the spend log cleanup job, on that job's schedule. Unset means rows are never deleted. " + "Set this well above health_check_interval because /health and the UI read the latest row per model." + ), + ) use_spend_logs_partitioning: bool | None = Field( None, description="If True and LiteLLM_SpendLogs has been converted to a range-partitioned table (db_scripts/partition_spend_logs.sql), retention cleanup drops expired partitions instead of deleting rows, and pre-creates upcoming partitions. Default is False.", @@ -3731,6 +3756,11 @@ class ProxyErrorTypes(str, enum.Enum): General authentication error """ + auth_provider_unavailable = "auth_provider_unavailable" + """ + The identity provider needed to authenticate the request (e.g. its JWKS endpoint) is unreachable + """ + internal_server_error = "internal_server_error" """ Internal server error @@ -3821,6 +3851,7 @@ class ProxyErrorTypes(str, enum.Enum): DB_CONNECTION_ERROR_TYPES: Final = ( httpx.ConnectError, + httpx.ConnectTimeout, httpx.ReadError, httpx.ReadTimeout, ) @@ -4499,6 +4530,9 @@ class JWTIssuerConfig(BaseModel): return self +DEFAULT_JWKS_STALE_TTL: Final = 3600 + + class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): """ A class to define the roles and permissions for a LiteLLM Proxy w/ JWT Auth. @@ -4514,6 +4548,8 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): - user_allowed_email_subdomain: If specified, only emails from specified subdomain will be allowed to access proxy. - end_user_id_jwt_field: The field in the JWT token that stores the end-user ID (maps to `LiteLLMEndUserTable`). Turn this off by setting to `None`. Enables end-user cost tracking. Use this for external customers. - public_key_ttl: Default - 600s. TTL for caching public JWT keys. + - public_key_stale_ttl: Default - 3600s. Extra time past `public_key_ttl` that the last-known-good JWKS response + stays usable while the identity provider is unreachable. Set to 0 to fail closed instead. - public_allowed_routes: list of allowed routes for authenticated but unknown litellm role jwt tokens. - enforce_rbac: If true, enforce RBAC for all routes. - custom_validate: A custom function to validates the JWT token. @@ -4564,6 +4600,15 @@ class LiteLLM_JWTAuth(LiteLLMPydanticObjectBase): user_id_upsert: bool = Field(default=False, description="If user doesn't exist, upsert them into the db.") end_user_id_jwt_field: str | None = None public_key_ttl: float = 600 + public_key_stale_ttl: float = Field( + default=DEFAULT_JWKS_STALE_TTL, + ge=0, + description=( + "Seconds beyond `public_key_ttl` that the last-known-good JWKS response stays usable while the identity " + "provider is unreachable. Bounds how long a signing key the provider has since removed can still be " + "trusted. Set to 0 to fail closed and reject requests as soon as the cached keys expire." + ), + ) public_allowed_routes: list[str] = ["public_routes"] enforce_rbac: bool = False roles_jwt_field: str | None = None # v2 on role mappings diff --git a/litellm/proxy/agent_endpoints/a2a_endpoints.py b/litellm/proxy/agent_endpoints/a2a_endpoints.py index cea30ffad52..bd02cfdf907 100644 --- a/litellm/proxy/agent_endpoints/a2a_endpoints.py +++ b/litellm/proxy/agent_endpoints/a2a_endpoints.py @@ -77,6 +77,27 @@ _PASCAL_TO_WIRE: Final[Mapping[str, str]] = { } +def _sse_event(payload: object) -> str: + """Frame a JSON-RPC object as a single A2A SSE event (``data: \\n\\n``).""" + return f"data: {json.dumps(payload)}\n\n" + + +def _to_jsonrpc_object(chunk: object) -> object: + """Coerce a streamed chunk to the JSON-RPC object it carries. + + Chunks arrive as SDK models, plain dicts, or, when a guardrail terminates a + stream, as an already serialized JSON-RPC object. + """ + if isinstance(chunk, (str, bytes, bytearray)): + try: + return json.loads(chunk) + except (json.JSONDecodeError, UnicodeDecodeError): + return chunk + if hasattr(chunk, "model_dump"): + return chunk.model_dump(mode="json", exclude_none=True) + return chunk + + def _build_message_send_params(params: dict[str, Any]) -> "MessageSendParams": """Build MessageSendParams from wire (0.3) or A2A 1.0 JSON-RPC params.""" from a2a.compat.v0_3.types import MessageSendParams @@ -280,6 +301,22 @@ async def _a2a_sse_event_source( await resp.aclose() +def _sse_streaming_response(generator: AsyncGenerator[str, None]) -> StreamingResponse: + # The upstream agent is only contacted once this generator is first pulled, so + # a slow first event leaves the response body idle for its whole + # time-to-first-token and an intermediary with an idle read timeout drops a + # healthy connection. Off until an operator sets an interval, and the + # buffering hint only goes out when there are keepalives to protect. + keepalive_interval: Final = coerce_keepalive_interval(litellm.sse_keepalive_ping_interval_seconds) + if keepalive_interval is None: + return StreamingResponse(generator, media_type="text/event-stream") + return StreamingResponse( + wrap_sse_stream_with_keepalive_pings(generator, keepalive_interval, ping_chunk=SSE_COMMENT_PING), + media_type="text/event-stream", + headers=_SSE_KEEPALIVE_HEADERS, + ) + + async def _forward_jsonrpc_sse( agent_url: str, body: Mapping[str, object], @@ -341,19 +378,7 @@ async def _forward_jsonrpc_sse( generator = _passthrough() - # The upstream agent is only contacted once this generator is first pulled, so - # a slow first event leaves the response body idle for its whole - # time-to-first-token and an intermediary with an idle read timeout drops a - # healthy connection. Off until an operator sets an interval, and the - # buffering hint only goes out when there are keepalives to protect. - keepalive_interval: Final = coerce_keepalive_interval(litellm.sse_keepalive_ping_interval_seconds) - if keepalive_interval is None: - return StreamingResponse(generator, media_type="text/event-stream") - return StreamingResponse( - wrap_sse_stream_with_keepalive_pings(generator, keepalive_interval, ping_chunk=SSE_COMMENT_PING), - media_type="text/event-stream", - headers=_SSE_KEEPALIVE_HEADERS, - ) + return _sse_streaming_response(generator) async def _handle_stream_message( @@ -373,9 +398,12 @@ async def _handle_stream_message( ) -> StreamingResponse: """Handle message/stream method via SDK functions. - When user_api_key_dict, request_data, and proxy_logging_obj are provided, - uses common_request_processing.async_streaming_data_generator with NDJSON - serializers so proxy hooks and cost injection apply. + The A2A JSON-RPC binding streams responses as SSE (text/event-stream) with + each JSON-RPC object framed as ``data: \n\n``, matching the official + a2a-sdk client which rejects any other Content-Type. When user_api_key_dict, + request_data, and proxy_logging_obj are provided, events are routed through + common_request_processing.async_streaming_data_generator so proxy hooks and + cost injection apply. """ from litellm.a2a_protocol import asend_message_streaming from litellm.a2a_protocol.main import A2A_SDK_AVAILABLE @@ -383,21 +411,18 @@ async def _handle_stream_message( if not A2A_SDK_AVAILABLE: async def _error_stream(): - yield ( - json.dumps( - { - "jsonrpc": "2.0", - "id": request_id, - "error": { - "code": -32603, - "message": "Server error: 'a2a' package not installed", - }, - } - ) - + "\n" + yield _sse_event( + { + "jsonrpc": "2.0", + "id": request_id, + "error": { + "code": -32603, + "message": "Server error: 'a2a' package not installed", + }, + } ) - return StreamingResponse(_error_stream(), media_type="application/x-ndjson") + return StreamingResponse(_error_stream(), media_type="text/event-stream") from a2a.compat.v0_3.types import SendStreamingMessageRequest @@ -409,18 +434,21 @@ async def _handle_stream_message( invalid_params_message: Final = f"Invalid params: {e}" async def _invalid_params_stream(): - yield ( - json.dumps( - { - "jsonrpc": "2.0", - "id": request_id, - "error": {"code": -32602, "message": invalid_params_message}, - } - ) - + "\n" + yield _sse_event( + { + "jsonrpc": "2.0", + "id": request_id, + "error": {"code": -32602, "message": invalid_params_message}, + } ) - return StreamingResponse(_invalid_params_stream(), media_type="application/x-ndjson") + return StreamingResponse(_invalid_params_stream(), media_type="text/event-stream") + + def _sse_chunk(chunk: object) -> str: + obj = _to_jsonrpc_object(chunk) + if isinstance(obj, dict): + obj = normalize_stream_event(obj, served_version, request_id=request_id) + return _sse_event(obj) async def stream_response(): try: @@ -448,32 +476,20 @@ async def _handle_stream_message( ProxyBaseLLMRequestProcessing, ) - def _ndjson_chunk(chunk: Any) -> str: - if hasattr(chunk, "model_dump"): - obj = chunk.model_dump(mode="json", exclude_none=True) - else: - obj = chunk - if isinstance(obj, dict): - obj = normalize_stream_event(obj, served_version, request_id=request_id) - return json.dumps(obj) + "\n" - - def _ndjson_error(proxy_exc: object) -> str: - return ( - json.dumps( - { - "jsonrpc": "2.0", - "id": request_id, - "error": { - "code": -32603, - "message": getattr( - proxy_exc, - "message", - f"Streaming error: {proxy_exc}", - ), - }, - } - ) - + "\n" + def _sse_error(proxy_exc: object) -> str: + return _sse_event( + { + "jsonrpc": "2.0", + "id": request_id, + "error": { + "code": -32603, + "message": getattr( + proxy_exc, + "message", + f"Streaming error: {proxy_exc}", + ), + }, + } ) async for line in ProxyBaseLLMRequestProcessing.async_streaming_data_generator( @@ -481,19 +497,13 @@ async def _handle_stream_message( user_api_key_dict=user_api_key_dict, request_data=request_data, proxy_logging_obj=proxy_logging_obj, - serialize_chunk=_ndjson_chunk, - serialize_error=_ndjson_error, + serialize_chunk=_sse_chunk, + serialize_error=_sse_error, ): yield line else: async for chunk in a2a_stream: - if hasattr(chunk, "model_dump"): - obj = chunk.model_dump(mode="json", exclude_none=True) - else: - obj = chunk - if isinstance(obj, dict): - obj = normalize_stream_event(obj, served_version, request_id=request_id) - yield json.dumps(obj) + "\n" + yield _sse_chunk(chunk) except Exception as e: verbose_proxy_logger.exception("Error streaming A2A response: %s", e) if ( @@ -511,21 +521,18 @@ async def _handle_stream_message( e = transformed_exception if isinstance(e, HTTPException): raise - yield ( - json.dumps( - { - "jsonrpc": "2.0", - "id": request_id, - "error": { - "code": -32603, - "message": f"Streaming error: {e}", - }, - } - ) - + "\n" + yield _sse_event( + { + "jsonrpc": "2.0", + "id": request_id, + "error": { + "code": -32603, + "message": f"Streaming error: {e}", + }, + } ) - return StreamingResponse(stream_response(), media_type="application/x-ndjson") + return _sse_streaming_response(stream_response()) @router.get( diff --git a/litellm/proxy/auth/auth_checks.py b/litellm/proxy/auth/auth_checks.py index 8708f96339f..e1cc91df1c7 100644 --- a/litellm/proxy/auth/auth_checks.py +++ b/litellm/proxy/auth/auth_checks.py @@ -247,14 +247,6 @@ def _raw_cache(cache: _RawCacheRead) -> _RawCacheRead: return cache -class _BudgetCacheRead(Protocol): - async def async_get_cache(self, *, key: str) -> "LiteLLM_BudgetTable | Mapping[str, object] | None": ... - - -def _budget_cache(cache: _BudgetCacheRead) -> _BudgetCacheRead: - return cache - - def _typed_request_body(request_body: dict) -> Mapping[str, object]: return request_body @@ -1190,33 +1182,35 @@ async def get_team_member_default_budget( cache_key: Final = f"team_member_default_budget:{budget_id}" - cached_budget: Final = await _budget_cache(user_api_key_cache).async_get_cache(key=cache_key) - if isinstance(cached_budget, LiteLLM_BudgetTable): + cached_budget: Final = await user_api_key_cache.async_get_cache( + key=cache_key, + model_type=LiteLLM_BudgetTable, + ) + if cached_budget is not None: return cached_budget - if isinstance(cached_budget, dict): - return LiteLLM_BudgetTable.model_validate(cached_budget) try: budget_record: Final = await _dictable_table(BudgetRepository(prisma_client)).find_unique( where={"budget_id": budget_id} ) - - if budget_record is None: - verbose_proxy_logger.warning("Team-default member budget not found in database: %s", budget_id) - return None - - await user_api_key_cache.async_set_cache( - key=cache_key, - value=budget_record.dict(), - ttl=get_management_object_ttl(user_api_key_cache), - ) - - return LiteLLM_BudgetTable.model_validate(budget_record.dict()) - except Exception: verbose_proxy_logger.exception("Error fetching team-default member budget %s", budget_id) return None + if budget_record is None: + verbose_proxy_logger.warning("Team-default member budget not found in database: %s", budget_id) + return None + + budget: Final = LiteLLM_BudgetTable.model_validate(budget_record.dict()) + await user_api_key_cache.async_set_cache( + key=cache_key, + value=budget, + model_type=LiteLLM_BudgetTable, + ttl=get_management_object_ttl(user_api_key_cache), + ) + + return budget + async def _apply_default_budget_to_end_user( end_user_obj: LiteLLM_EndUserTable, diff --git a/litellm/proxy/auth/auth_exception_handler.py b/litellm/proxy/auth/auth_exception_handler.py index 603e72463bc..233679126f8 100644 --- a/litellm/proxy/auth/auth_exception_handler.py +++ b/litellm/proxy/auth/auth_exception_handler.py @@ -2,12 +2,14 @@ Handles Authentication Errors """ +from collections.abc import Mapping from typing import TYPE_CHECKING, Any, Final from fastapi import HTTPException, Request, status import litellm from litellm._logging import verbose_proxy_logger +from litellm.constants import EMPTY_MAPPING from litellm.integrations.otel.runtime import seed_request_identity from litellm.proxy._types import ( LitellmUserRoles, @@ -33,12 +35,25 @@ else: Span = Any +def _with_requester_ip_address(request_data: dict[str, object], requester_ip: str | None) -> dict[str, object]: + """Auth gate rejections are raised before `add_litellm_data_to_request` records the + caller IP, so their failure logs would otherwise carry no IP nor key/user identity.""" + if not requester_ip: + return request_data + key: Final = "litellm_metadata" if "litellm_metadata" in request_data else "metadata" + metadata: Final = request_data.get(key) + base: Final[Mapping[str, object]] = metadata if isinstance(metadata, Mapping) else EMPTY_MAPPING + if base.get("requester_ip_address"): + return request_data + return {**request_data, key: {**base, "requester_ip_address": requester_ip}} # mutable-ok: logging needs dicts + + class UserAPIKeyAuthExceptionHandler: @staticmethod async def _handle_authentication_error( e: Exception, request: Request, - request_data: dict, + request_data: dict[str, object], route: str, parent_otel_span: Span | None, api_key: str, @@ -92,7 +107,7 @@ class UserAPIKeyAuthExceptionHandler: # raise the exception to the caller requester_ip: Final = _get_request_ip_address( request=request, - use_x_forwarded_for=general_settings.get("use_x_forwarded_for", False), + use_x_forwarded_for=general_settings.get("use_x_forwarded_for") is True, ) verbose_proxy_logger.exception( "litellm.proxy.proxy_server.user_api_key_auth(): Exception occured - %s\nRequester IP Address:%s", @@ -129,11 +144,14 @@ class UserAPIKeyAuthExceptionHandler: resolve_llm_provider_for_rate_limit, ) - _, e.llm_provider = resolve_llm_provider_for_rate_limit(request_data.get("model")) + budget_model: Final = request_data.get("model") + _, e.llm_provider = resolve_llm_provider_for_rate_limit( + budget_model if isinstance(budget_model, str) else None + ) # Allow callbacks to transform the error response transformed_exception: Final = await proxy_logging_obj.post_call_failure_hook( - request_data=request_data, + request_data=_with_requester_ip_address(request_data, requester_ip), original_exception=e, user_api_key_dict=user_api_key_dict, error_type=ProxyErrorTypes.auth_error, diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 912a0b0ebd0..ce662ee0374 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -311,6 +311,12 @@ _BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = ( # the request away from the admin's pinned configuration. "nvcf_function_id", "use_ssl", + # Per-deployment opt-in that hands the whole call to the Rust core. It is a + # deployment decision, not a request one: the Rust path uses its own client + # rather than the one the deployment configured, and reports no post_call, + # so a caller-supplied value picks a transport and a callback surface the + # admin did not choose. + "rust", # SDK-only field; also rejected outright in is_request_body_safe. "model_list", "vertex_ai_credentials", diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 1e3265af967..39e6ca9a369 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -8,12 +8,16 @@ JWT token must have 'litellm_proxy_admin' in scope. from __future__ import annotations +import asyncio import fnmatch import hashlib import os import re -from typing import Any, Final, Literal, NoReturn, cast +import time +from collections.abc import Awaitable, Callable +from typing import Any, Final, Literal, NoReturn, TypeVar, cast +import httpx import jwt from cryptography import x509 from cryptography.hazmat.backends import default_backend @@ -25,6 +29,7 @@ from litellm._logging import verbose_proxy_logger from litellm.litellm_core_utils.dot_notation_indexing import get_nested_value from litellm.llms.custom_httpx.httpx_handler import HTTPHandler from litellm.proxy._types import ( + DEFAULT_JWKS_STALE_TTL, RBAC_ROLES, JWKKeyValue, JWTAuthBuilderResult, @@ -74,6 +79,32 @@ class NoMatchingJWTPublicKeyError(Exception): """Raised when a JWKS endpoint returns no key matching the requested ``kid``.""" +class JWKSUnreachableError(Exception): + """Raised when an IdP's JWKS / OIDC discovery endpoint is unreachable and no cached copy is left to fall back on.""" + + +JWKS_FETCH_ATTEMPTS: Final = 3 +JWKS_FETCH_RETRY_BACKOFF_SECONDS: Final = 0.25 +JWKS_UNREACHABLE_BACKOFF_SECONDS: Final = 30 +STALE_CACHE_KEY_PREFIX: Final = "litellm_stale_" +STALE_WRITTEN_AT_CACHE_KEY_PREFIX: Final = "litellm_stale_written_at_" +UNREACHABLE_CACHE_KEY_PREFIX: Final = "litellm_jwks_unreachable_" + +_CachedValueT = TypeVar("_CachedValueT", bound=JWKKeyValue | str) + + +def jwks_unavailable_exception(error: JWKSUnreachableError) -> ProxyException: + return ProxyException( + message=( + "Service Unavailable, the identity provider's JWKS endpoint is temporarily " + f"unreachable, so the JWT signature could not be verified. Please retry shortly. Error: {error}" + ), + type=ProxyErrorTypes.auth_provider_unavailable, + param="None", + code=status.HTTP_503_SERVICE_UNAVAILABLE, + ) + + class JWTHandler: """ - treat the sub id passed in as the user id @@ -121,6 +152,8 @@ class JWTHandler: ) -> None: self.http_handler = HTTPHandler() self.leeway = 0 + # Per-cache-key locks so a TTL lapse triggers one refresh instead of one per in-flight request. + self._refresh_locks: dict[str, asyncio.Lock] = {} # mutable-ok: lock registry, keyed by JWKS url def update_environment( self, @@ -611,13 +644,151 @@ class JWTHandler: if ".well-known/openid-configuration" not in url: return url - cache_key: Final = f"litellm_oidc_discovery_{url}" - cached_jwks_uri: Final = await self.user_api_key_cache.async_get_cache(cache_key) - if cached_jwks_uri is not None: - return cached_jwks_uri + return await self._cached_with_stale_fallback( + cache_key=f"litellm_oidc_discovery_{url}", + ttl=self._get_public_key_cache_ttl(), + refresh=lambda: self._fetch_jwks_uri_from_discovery(url), + log_context="an OIDC discovery lookup", + ) + async def _get_with_transient_retries(self, url: str) -> httpx.Response: + """GET ``url``, retrying transport failures so one IdP blip does not fail the request.""" + for attempt in range(1, JWKS_FETCH_ATTEMPTS): + try: + return await self.http_handler.get(url) + except httpx.TransportError as e: + verbose_proxy_logger.warning( + "JWT Auth: %s fetching %s (attempt %s/%s), retrying: %s", + type(e).__name__, + url, + attempt, + JWKS_FETCH_ATTEMPTS, + e, + ) + await asyncio.sleep(JWKS_FETCH_RETRY_BACKOFF_SECONDS * attempt) + + try: + return await self.http_handler.get(url) + except httpx.TransportError as e: + raise JWKSUnreachableError(f"{type(e).__name__} fetching {url} after {JWKS_FETCH_ATTEMPTS} attempts") from e + + async def _get_cached_value(self, cache_key: str) -> _CachedValueT | None: + cached: Final = await self.user_api_key_cache.async_get_cache(cache_key) + return cast("_CachedValueT | None", cached) # cast-ok: cache reads are untyped + + async def _get_cached_timestamp(self, cache_key: str) -> float | None: + cached: Final = await self.user_api_key_cache.async_get_cache(cache_key) + # A JSON round-trip through Redis hands a whole-number epoch back as an int. + return float(cached) if isinstance(cached, (int, float)) else None + + async def _put_cached_value(self, cache_key: str, value: JWKKeyValue | str | float, ttl: float) -> None: + await self.user_api_key_cache.async_set_cache(key=cache_key, value=value, ttl=ttl) + + async def _cached_with_stale_fallback( + self, + cache_key: str, + ttl: float, + refresh: Callable[[], Awaitable[_CachedValueT]], + log_context: str, + ) -> _CachedValueT: + """Read ``cache_key``, refreshing it through a single-flight lock on a miss.""" + cached: Final[_CachedValueT | None] = await self._get_cached_value(cache_key) + if cached is not None: + return cached + + lock: Final = self._refresh_locks.setdefault(cache_key, asyncio.Lock()) + async with lock: + cached_after_lock: Final[_CachedValueT | None] = await self._get_cached_value(cache_key) + if cached_after_lock is not None: + return cached_after_lock + return await self._refresh_or_serve_stale( + cache_key=cache_key, ttl=ttl, refresh=refresh, log_context=log_context + ) + + async def _refresh_or_serve_stale( + self, + cache_key: str, + ttl: float, + refresh: Callable[[], Awaitable[_CachedValueT]], + log_context: str, + ) -> _CachedValueT: + """Refresh ``cache_key`` from the IdP, falling back to the last-known-good copy when it is unreachable. + + Signing keys rotate rarely, so a last-known-good key beats failing authentication during an IdP blip. + How long a key the IdP has since removed stays trusted is bounded by ``public_key_ttl`` + + ``public_key_stale_ttl`` measured from when the copy was taken, and that bound is enforced here on every + read rather than baked into the cache entry's own expiry. An operator who lowers ``public_key_stale_ttl``, + or sets it to 0 to fail closed, is usually doing it mid-incident, and a copy written under the old longer + setting would otherwise stay servable until it aged out on its own. A copy whose write time cannot be + established is not servable, so the bound cannot be dodged by losing the timestamp. + """ + stale_ttl: Final = self._get_public_key_stale_ttl() + outcome: Final = await self._refresh_or_record_outage( + cache_key=cache_key, ttl=ttl, stale_ttl=stale_ttl, refresh=refresh + ) + if not isinstance(outcome, JWKSUnreachableError): + return outcome + if stale_ttl <= 0: + raise outcome + + stale: Final[_CachedValueT | None] = await self._get_cached_value(f"{STALE_CACHE_KEY_PREFIX}{cache_key}") + age: Final = await self._stale_copy_age(cache_key) + lifetime: Final = ttl + stale_ttl + if stale is None or age is None or age > lifetime: + raise outcome + verbose_proxy_logger.warning( + "JWT Auth: identity provider unreachable, authenticating %s against a stale JWKS copy of %s " + "(last refreshed %.0fs ago, stops being trusted in %.0fs). Refresh failed: %s", + log_context, + cache_key, + age, + max(lifetime - age, 0), + outcome, + ) + return stale + + async def _stale_copy_age(self, cache_key: str) -> float | None: + written_at: Final = await self._get_cached_timestamp(f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{cache_key}") + return None if written_at is None else time.time() - written_at + + async def _refresh_or_record_outage( + self, + cache_key: str, + ttl: float, + stale_ttl: float, + refresh: Callable[[], Awaitable[_CachedValueT]], + ) -> _CachedValueT | JWKSUnreachableError: + """Refresh ``cache_key``, returning the outage as a value rather than raising it. + + A failed refresh is remembered for ``JWKS_UNREACHABLE_BACKOFF_SECONDS`` so a sustained outage costs one + fetch per window instead of one per request serialised behind the refresh lock. + """ + unreachable_cache_key: Final = f"{UNREACHABLE_CACHE_KEY_PREFIX}{cache_key}" + recent_failure: Final[str | None] = await self._get_cached_value(unreachable_cache_key) + if recent_failure is not None: + return JWKSUnreachableError(recent_failure) + + try: + refreshed: Final = await refresh() + except JWKSUnreachableError as e: + await self._put_cached_value( + cache_key=unreachable_cache_key, value=str(e), ttl=JWKS_UNREACHABLE_BACKOFF_SECONDS + ) + return e + + await self._put_cached_value(cache_key=cache_key, value=refreshed, ttl=ttl) + if stale_ttl > 0: + await self._put_cached_value( + cache_key=f"{STALE_CACHE_KEY_PREFIX}{cache_key}", value=refreshed, ttl=ttl + stale_ttl + ) + await self._put_cached_value( + cache_key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{cache_key}", value=time.time(), ttl=ttl + stale_ttl + ) + return refreshed + + async def _fetch_jwks_uri_from_discovery(self, url: str) -> str: verbose_proxy_logger.debug("JWT Auth: Fetching OIDC discovery document from %s", url) - response: Final = await self.http_handler.get(url) + response: Final = await self._get_with_transient_retries(url) if response.status_code != 200: raise Exception( f"JWT Auth: OIDC discovery endpoint {url} returned status {response.status_code}: {response.text}" @@ -632,11 +803,6 @@ class JWTHandler: raise Exception(f"JWT Auth: OIDC discovery document at {url} does not contain a 'jwks_uri' field.") verbose_proxy_logger.debug("JWT Auth: Resolved OIDC discovery %s -> jwks_uri=%s", url, jwks_uri) - await self.user_api_key_cache.async_set_cache( - key=cache_key, - value=jwks_uri, - ttl=self._get_public_key_cache_ttl(), - ) return jwks_uri def _get_public_key_cache_ttl(self) -> float: @@ -645,33 +811,36 @@ class JWTHandler: return 600 return litellm_jwtauth.public_key_ttl + def _get_public_key_stale_ttl(self) -> float: + litellm_jwtauth: Final = getattr(self, "litellm_jwtauth", None) + if litellm_jwtauth is None: + return DEFAULT_JWKS_STALE_TTL + return litellm_jwtauth.public_key_stale_ttl + + async def _fetch_jwks_keys(self, resolved_jwks_url: str) -> JWKKeyValue: + response: Final = await self._get_with_transient_retries(resolved_jwks_url) + if response.status_code != 200: + raise Exception( + f"JWT Auth: JWKS endpoint {resolved_jwks_url} returned status {response.status_code}: {response.text}" + ) + + try: + response_json: Final = response.json() + except Exception as e: + verbose_proxy_logger.error("Error parsing response: %s. Original Response: %s", e, response.text) + raise Exception(f"Error parsing response: {e}. Check server logs for original response.") + + keys: Final = response_json["keys"] if "keys" in response_json else response_json + return cast(JWKKeyValue, keys) # cast-ok: JWTKeyItem declares only `kid`, validating would drop key material + async def _get_public_key_from_jwks_url(self, jwks_url: str, kid: str | None) -> dict: resolved_jwks_url: Final = await self._resolve_jwks_url(jwks_url) - cache_key: Final = f"litellm_jwt_auth_keys_{resolved_jwks_url}" - - cached_keys: Final = await self.user_api_key_cache.async_get_cache(cache_key) - - if cached_keys is None: - response: Final = await self.http_handler.get(resolved_jwks_url) - - try: - response_json: Final = response.json() - except Exception as e: - verbose_proxy_logger.error("Error parsing response: %s. Original Response: %s", e, response.text) - raise Exception(f"Error parsing response: {e}. Check server logs for original response.") - - if "keys" in response_json: - keys: JWKKeyValue = response_json["keys"] - else: - keys = response_json - - await self.user_api_key_cache.async_set_cache( - key=cache_key, - value=keys, - ttl=self._get_public_key_cache_ttl(), - ) - else: - keys = cached_keys + keys: Final = await self._cached_with_stale_fallback( + cache_key=f"litellm_jwt_auth_keys_{resolved_jwks_url}", + ttl=self._get_public_key_cache_ttl(), + refresh=lambda: self._fetch_jwks_keys(resolved_jwks_url), + log_context=f"kid={kid}", + ) public_key: Final = self.parse_keys(keys=keys, kid=kid) if public_key is not None: @@ -692,6 +861,9 @@ class JWTHandler: return await self._get_public_key_from_jwks_url(jwks_url=key_url, kid=kid) except NoMatchingJWTPublicKeyError as e: verbose_proxy_logger.debug("JWT Auth: No matching public key found at %s: %s", key_url, e) + except JWKSUnreachableError as e: + verbose_proxy_logger.error("JWT Auth: JWKS endpoint %s unreachable: %s", key_url, e) + raise jwks_unavailable_exception(e) from e raise NoMatchingJWTPublicKeyError(f"No matching public key found. keys={keys_url_list}, kid={kid}") @@ -969,10 +1141,14 @@ class JWTHandler: ) async def _auth_jwt_with_issuer(self, token: str, issuer_config: JWTIssuerConfig, kid: str | None) -> dict: - public_key: Final = await self._get_public_key_from_jwks_url( - jwks_url=self._get_jwks_url_for_issuer(issuer_config=issuer_config), - kid=kid, - ) + try: + public_key: Final = await self._get_public_key_from_jwks_url( + jwks_url=self._get_jwks_url_for_issuer(issuer_config=issuer_config), + kid=kid, + ) + except JWKSUnreachableError as e: + raise jwks_unavailable_exception(e) from e + try: payload: Final = self._decode_jwt_with_public_key( token=token, diff --git a/litellm/proxy/auth/route_checks.py b/litellm/proxy/auth/route_checks.py index 04eb7ab326b..cea21ca088b 100644 --- a/litellm/proxy/auth/route_checks.py +++ b/litellm/proxy/auth/route_checks.py @@ -1,4 +1,5 @@ import re +from collections.abc import Sequence from typing import Final from fastapi import HTTPException, Request, status @@ -165,6 +166,19 @@ class RouteChecks: if RouteChecks._is_get_mcp_server_discovery_route(route=route, request=request): return True + # Agent registry CRUD moved from llm_api_routes into + # management_routes so DISABLE_LLM_API_ENDPOINTS stops + # blocking it. Keys configured with + # allowed_routes=["llm_api_routes"] before that split + # could reach these paths, so keep them reachable here; + # the handlers in agent_endpoints/endpoints.py still + # enforce proxy-admin on writes and scope reads by role. + if RouteChecks.check_route_access( + route=route, + allowed_routes=LiteLLMRoutes.agent_management_routes.value, + ): + return True + # check if wildcard pattern is allowed for allowed_route in valid_token.allowed_routes: if RouteChecks._route_matches_wildcard_pattern(route=route, pattern=allowed_route): @@ -367,7 +381,7 @@ class RouteChecks: if RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.mcp_inference_routes.value): return True - if RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.agent_routes.value): + if RouteChecks.check_route_access(route=route, allowed_routes=LiteLLMRoutes.agent_inference_routes.value): return True if route in LiteLLMRoutes.litellm_native_routes.value: @@ -558,13 +572,13 @@ class RouteChecks: return False @staticmethod - def check_route_access(route: str, allowed_routes: list[str]) -> bool: + def check_route_access(route: str, allowed_routes: Sequence[str]) -> bool: """ Check if a route has access by checking both exact matches and patterns Args: route (str): The route to check - allowed_routes (list): List of allowed routes/patterns + allowed_routes (Sequence): Allowed routes/patterns Returns: bool: True if route is allowed, False otherwise @@ -579,10 +593,12 @@ class RouteChecks: # wildcard match route is in allowed_routes # e.g calling /anthropic/v1/messages is allowed if allowed_routes has /anthropic/* ######################################################### - wildcard_allowed_routes = [route for route in allowed_routes if RouteChecks._is_wildcard_pattern(pattern=route)] - for allowed_route in wildcard_allowed_routes: - if RouteChecks._route_matches_wildcard_pattern(route=route, pattern=allowed_route): - return True + if any( + RouteChecks._route_matches_wildcard_pattern(route=route, pattern=allowed_route) + for allowed_route in allowed_routes + if RouteChecks._is_wildcard_pattern(pattern=allowed_route) + ): + return True ######################################################### # pattern match route is in allowed_routes diff --git a/litellm/proxy/common_utils/reset_budget_job.py b/litellm/proxy/common_utils/reset_budget_job.py index 7b7cba5fc42..b28b7291a4c 100644 --- a/litellm/proxy/common_utils/reset_budget_job.py +++ b/litellm/proxy/common_utils/reset_budget_job.py @@ -17,6 +17,7 @@ from litellm.constants import ( RESET_BUDGET_JOB_MAX_CHUNKS_PER_RUN, ) from litellm.proxy._types import ( + DB_RETRY_SAFE_ERROR_TYPES, LiteLLM_BudgetTableFull, LiteLLM_EndUserTable, LiteLLM_TeamTable, @@ -29,6 +30,7 @@ from litellm.proxy.common_utils.timezone_utils import ( get_budget_reset_settings, ) from litellm.proxy.common_utils.user_api_key_cache import tag_cache_key +from litellm.proxy.db.exception_handler import call_with_db_reconnect_retry from litellm.proxy.utils import PrismaClient, ProxyLogging from litellm.repositories.organization_repository import OrganizationRepository from litellm.repositories.prisma_protocols import ReadOnlyTable, SpendLinkedTable @@ -236,6 +238,24 @@ class ResetBudgetJob: await self.reset_budget_for_litellm_budget_table() await self.reset_budget_windows() + async def _with_db_retry(self, operation: Callable[[], Awaitable[_RowT]], *, reason: str) -> _RowT: + """Reconnect and retry once on a transport error, so a dropped connection + costs one retry instead of the whole tick. + """ + return await call_with_db_reconnect_retry(self.prisma_client, operation, reason=reason) + + async def _with_db_write_retry(self, operation: Callable[[], Awaitable[_RowT]], *, reason: str) -> _RowT: + """Same, for writes: only replay when the statements provably never + reached the database. A reset zeroes spend unconditionally, so replaying + an ambiguous commit would erase spend accrued since it landed. + """ + return await call_with_db_reconnect_retry( + self.prisma_client, + operation, + reason=reason, + retry_safe_error_types=DB_RETRY_SAFE_ERROR_TYPES, + ) + @staticmethod async def _invalidate_spend_counter(counter_key: str) -> None: """Zero a spend counter so a DB-row reset takes effect immediately. @@ -301,16 +321,24 @@ class ResetBudgetJob: """Read the rows the cascade will zero, so their counters can be invalidated once the transaction commits.""" try: - return tuple(await table.find_many(where=where)) + return tuple( + await self._with_db_retry( + lambda: table.find_many(where=where), + reason=f"reset_budget_read_{log_subject.replace(' ', '_')}_failure", + ) + ) except Exception as e: verbose_proxy_logger.warning("Failed to fetch %s for counter invalidation: %s", log_subject, e) return () async def _collect_endusers_to_reset(self, budget_ids: Sequence[str]) -> tuple[_EndUserRow, ...]: - linked: Final[Sequence[_EndUserRow] | None] = await self.prisma_client.get_data( - table_name="enduser", - query_type="find_all", - budget_id_list=list(budget_ids), + linked: Final[Sequence[_EndUserRow] | None] = await self._with_db_retry( + lambda: self.prisma_client.get_data( + table_name="enduser", + query_type="find_all", + budget_id_list=list(budget_ids), + ), + reason="reset_budget_read_endusers_failure", ) if litellm.max_end_user_budget_id is None or litellm.max_end_user_budget_id not in budget_ids: return tuple(linked or ()) @@ -384,6 +412,12 @@ class ResetBudgetJob: if not cascade.budget_ids: return + await self._with_db_write_retry( + lambda: self._commit_budget_cascade_once(cascade), + reason="reset_budget_write_budget_cascade_failure", + ) + + async def _commit_budget_cascade_once(self, cascade: _BudgetCascade) -> None: enduser_ids: Final = tuple(row.user_id for row in cascade.endusers) async with budget_cascade_unit_of_work(self.prisma_client.db.batch_) as uow: uow.team_memberships.queue_spend_zero(where=_budget_link_where(cascade.budget_ids)) @@ -404,11 +438,14 @@ class ResetBudgetJob: async def _reset_expired_budget_cascade(self) -> _BudgetCascadeCommitted | _BudgetCascadeFailed: now: Final = datetime.now(timezone.utc) try: - budgets_to_reset: Final[Sequence[LiteLLM_BudgetTableFull] | None] = await self.prisma_client.get_data( - table_name="budget", - query_type="find_all", - reset_at=now, - limit=RESET_BUDGET_JOB_BATCH_SIZE, + budgets_to_reset: Final[Sequence[LiteLLM_BudgetTableFull] | None] = await self._with_db_retry( + lambda: self.prisma_client.get_data( + table_name="budget", + query_type="find_all", + reset_at=now, + limit=RESET_BUDGET_JOB_BATCH_SIZE, + ), + reason="reset_budget_read_budgets_failure", ) cascade: Final = await self._collect_budget_cascade(budgets_to_reset or ()) except Exception as e: @@ -492,11 +529,14 @@ class ResetBudgetJob: in-memory during auth checks. """ table: Final[ReadOnlyTable] = EndUserRepository(self.prisma_client).table - rows: Final = await table.find_many( - where={ - "budget_id": None, - "spend": {"gt": 0}, - }, + rows: Final = await self._with_db_retry( + lambda: table.find_many( + where={ + "budget_id": None, + "spend": {"gt": 0}, + }, + ), + reason="reset_budget_read_endusers_without_budget_id_failure", ) return [LiteLLM_EndUserTable.model_validate(row.dict()) for row in rows] @@ -511,6 +551,12 @@ class ResetBudgetJob: aborts the entire batch — silently leaving spend over the cap and budget_reset_at unchanged forever. """ + await self._with_db_write_retry( + lambda: self._write_key_reset_updates_once(updated_keys), + reason="reset_budget_write_keys_failure", + ) + + async def _write_key_reset_updates_once(self, updated_keys: list[LiteLLM_VerificationToken]) -> None: async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow: for k in updated_keys: if k.token is None: @@ -525,6 +571,12 @@ class ResetBudgetJob: that trips Prisma's DataError on rows carrying unrecognised fields (see #27730). """ + await self._with_db_write_retry( + lambda: self._write_user_reset_updates_once(updated_users), + reason="reset_budget_write_users_failure", + ) + + async def _write_user_reset_updates_once(self, updated_users: list[LiteLLM_UserTable]) -> None: async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow: for u in updated_users: uow.users.queue_spend_reset(user_id=u.user_id, budget_reset_at=u.budget_reset_at) @@ -537,6 +589,12 @@ class ResetBudgetJob: that trips Prisma's DataError on rows carrying unrecognised fields (see #27730). """ + await self._with_db_write_retry( + lambda: self._write_team_reset_updates_once(updated_teams), + reason="reset_budget_write_teams_failure", + ) + + async def _write_team_reset_updates_once(self, updated_teams: list[LiteLLM_TeamTable]) -> None: async with spend_reset_unit_of_work(self.prisma_client.db.batch_) as uow: for t in updated_teams: uow.teams.queue_spend_reset(team_id=t.team_id, budget_reset_at=t.budget_reset_at) @@ -579,12 +637,15 @@ class ResetBudgetJob: start_time: Final = time.time() keys_to_reset: list[LiteLLM_VerificationToken] | None = None try: - keys_to_reset = await self.prisma_client.get_data( - table_name="key", - query_type="find_all", - expires=now, - reset_at=now, - limit=RESET_BUDGET_JOB_BATCH_SIZE, + keys_to_reset = await self._with_db_retry( + lambda: self.prisma_client.get_data( + table_name="key", + query_type="find_all", + expires=now, + reset_at=now, + limit=RESET_BUDGET_JOB_BATCH_SIZE, + ), + reason="reset_budget_read_keys_failure", ) verbose_proxy_logger.debug("Keys to reset %s", json.dumps(keys_to_reset, indent=4, default=str)) updated_keys: Final[list[LiteLLM_VerificationToken]] = [] @@ -684,11 +745,14 @@ class ResetBudgetJob: start_time: Final = time.time() users_to_reset: list[LiteLLM_UserTable] | None = None try: - users_to_reset = await self.prisma_client.get_data( - table_name="user", - query_type="find_all", - reset_at=now, - limit=RESET_BUDGET_JOB_BATCH_SIZE, + users_to_reset = await self._with_db_retry( + lambda: self.prisma_client.get_data( + table_name="user", + query_type="find_all", + reset_at=now, + limit=RESET_BUDGET_JOB_BATCH_SIZE, + ), + reason="reset_budget_read_users_failure", ) updated_users: Final[list[LiteLLM_UserTable]] = [] failed_users: Final = [] @@ -795,11 +859,14 @@ class ResetBudgetJob: start_time: Final = time.time() teams_to_reset: list[LiteLLM_TeamTable] | None = None try: - teams_to_reset = await self.prisma_client.get_data( - table_name="team", - query_type="find_all", - reset_at=now, - limit=RESET_BUDGET_JOB_BATCH_SIZE, + teams_to_reset = await self._with_db_retry( + lambda: self.prisma_client.get_data( + table_name="team", + query_type="find_all", + reset_at=now, + limit=RESET_BUDGET_JOB_BATCH_SIZE, + ), + reason="reset_budget_read_teams_failure", ) updated_teams: Final[list[LiteLLM_TeamTable]] = [] failed_teams: Final = [] @@ -937,8 +1004,11 @@ class ResetBudgetJob: # --- Keys --- try: - key_rows: Final = await self.prisma_client.db.query_raw( - 'SELECT token, budget_limits FROM "LiteLLM_VerificationToken" WHERE budget_limits IS NOT NULL' + key_rows: Final = await self._with_db_retry( + lambda: self.prisma_client.db.query_raw( + 'SELECT token, budget_limits FROM "LiteLLM_VerificationToken" WHERE budget_limits IS NOT NULL' + ), + reason="reset_budget_read_key_windows_failure", ) for row in key_rows: raw = row["budget_limits"] @@ -957,17 +1027,23 @@ class ResetBudgetJob: ): changed = True if changed: - await VerificationTokenRepository(self.prisma_client).table.update( - where={"token": row["token"]}, - data={"budget_limits": json.dumps(windows)}, + await self._with_db_write_retry( + lambda: VerificationTokenRepository(self.prisma_client).table.update( + where={"token": row["token"]}, + data={"budget_limits": json.dumps(windows)}, + ), + reason="reset_budget_write_key_windows_failure", ) except Exception as e: verbose_proxy_logger.exception("Failed to reset budget windows for keys: %s", e) # --- Teams --- try: - team_rows: Final = await self.prisma_client.db.query_raw( - 'SELECT team_id, budget_limits FROM "LiteLLM_TeamTable" WHERE budget_limits IS NOT NULL' + team_rows: Final = await self._with_db_retry( + lambda: self.prisma_client.db.query_raw( + 'SELECT team_id, budget_limits FROM "LiteLLM_TeamTable" WHERE budget_limits IS NOT NULL' + ), + reason="reset_budget_read_team_windows_failure", ) for row in team_rows: raw = row["budget_limits"] @@ -986,9 +1062,12 @@ class ResetBudgetJob: ): changed = True if changed: - await TeamRepository(self.prisma_client).table.update( - where={"team_id": row["team_id"]}, - data={"budget_limits": json.dumps(windows)}, + await self._with_db_write_retry( + lambda: TeamRepository(self.prisma_client).table.update( + where={"team_id": row["team_id"]}, + data={"budget_limits": json.dumps(windows)}, + ), + reason="reset_budget_write_team_windows_failure", ) except Exception as e: verbose_proxy_logger.exception("Failed to reset budget windows for teams: %s", e) diff --git a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py index d19023862cb..e97e9f6e683 100644 --- a/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py +++ b/litellm/proxy/db/db_transaction_queue/spend_log_cleanup.py @@ -222,6 +222,21 @@ class SpendLogCleanup: remaining_ms: Final = int((deadline - time.monotonic()) * 1000) return max(1, min(int(self.batch_timeout_seconds * 1000), remaining_ms)) + @staticmethod + def _group_deadline(overall_deadline: float, groups_remaining: int) -> float: + """ + Give each pending cleanup group an equal share of the time left. + + A single group keeps the whole run budget, while a persistent backlog + on an earlier group cannot starve a later group. + """ + if groups_remaining == 1: + return overall_deadline + current_time: Final = time.monotonic() + if current_time >= overall_deadline: + return overall_deadline + return current_time + (overall_deadline - current_time) / groups_remaining + def _remaining_timeout_ms(self, deadline: float) -> RemainingTimeoutMs: """ The per-statement bound for work this job delegates, as a callable. @@ -477,6 +492,18 @@ class SpendLogCleanup: deadline=deadline, ) + async def _delete_old_health_check_rows( + self, prisma_client: PrismaClient, cutoff_date: datetime, deadline: float + ) -> TableCleanupResult: + return await self._delete_old_rows_batched( + prisma_client, + cutoff_date, + table_name="LiteLLM_HealthCheckTable", + key_columns=("health_check_id",), + time_column="checked_at", + deadline=deadline, + ) + async def _clean_spend_log_tables( self, prisma_client: PrismaClient, deadline: float ) -> tuple[TableCleanupResult, ...]: @@ -526,6 +553,19 @@ class SpendLogCleanup: verbose_proxy_logger.info("Deleted %s expired auto-router session rollup rows", sessions_result.rows_deleted) return (sessions_result,) + async def _clean_health_checks( + self, prisma_client: PrismaClient, retention_seconds: int, deadline: float + ) -> tuple[TableCleanupResult, ...]: + health_check_cutoff: Final = datetime.now(timezone.utc) - timedelta(seconds=float(retention_seconds)) + health_checks_result: Final = await self._delete_old_health_check_rows( + prisma_client, health_check_cutoff, deadline + ) + verbose_proxy_logger.info( + "Deleted %s expired health-check rows", + health_checks_result.rows_deleted, + ) + return (health_checks_result,) + @staticmethod def _run_outcome(results: tuple[TableCleanupResult, ...]) -> RunOutcome: """ @@ -558,7 +598,12 @@ class SpendLogCleanup: autorouter_retention_seconds: Final = self._retention_seconds_for( "maximum_autorouter_session_retention_period" ) - if not delete_spend_logs and autorouter_retention_seconds is None: + health_check_retention_seconds: Final = self._retention_seconds_for("maximum_health_check_retention_period") + if ( + not delete_spend_logs + and autorouter_retention_seconds is None + and health_check_retention_seconds is None + ): SpendLogCleanupMetrics.record_run("skipped_disabled") return @@ -585,19 +630,45 @@ class SpendLogCleanup: return deadline: Final = time.monotonic() + self.run_budget_seconds + configured_group_count: Final = ( + int(delete_spend_logs and self.retention_seconds is not None) + + int(autorouter_retention_seconds is not None) + + int(health_check_retention_seconds is not None) + ) spend_log_results: Final = ( - await self._clean_spend_log_tables(prisma_client, deadline) + await self._clean_spend_log_tables( + prisma_client, + self._group_deadline(deadline, configured_group_count), + ) if delete_spend_logs and self.retention_seconds is not None else () ) + remaining_groups_after_spend_logs: Final = int(autorouter_retention_seconds is not None) + int( + health_check_retention_seconds is not None + ) session_results: Final = ( - await self._clean_session_rollup(prisma_client, autorouter_retention_seconds, deadline) + await self._clean_session_rollup( + prisma_client, + autorouter_retention_seconds, + self._group_deadline(deadline, remaining_groups_after_spend_logs), + ) if autorouter_retention_seconds is not None else () ) + health_check_results: Final = ( + await self._clean_health_checks( + prisma_client, + health_check_retention_seconds, + deadline, + ) + if health_check_retention_seconds is not None + else () + ) - SpendLogCleanupMetrics.record_run(self._run_outcome(spend_log_results + session_results)) + SpendLogCleanupMetrics.record_run( + self._run_outcome(spend_log_results + session_results + health_check_results) + ) except Exception as e: # .exception() captures the traceback; str(e) alone on a Prisma/DB diff --git a/litellm/proxy/db/db_url_settings.py b/litellm/proxy/db/db_url_settings.py index 17e631995cd..d393aa1b977 100644 --- a/litellm/proxy/db/db_url_settings.py +++ b/litellm/proxy/db/db_url_settings.py @@ -11,10 +11,12 @@ The env var names this module reads are exactly the ones emitted by the (``helm/litellm/templates/_helpers.tpl``). Both auth styles and both endpoints are covered: - * IAM auth (``IAM_TOKEN_DB_AUTH`` truthy): mint a short-lived RDS IAM - token and embed it as the password. The writer URL is always - (re)written because the token is freshly minted on every startup. The - chart omits ``DATABASE_PASSWORD`` in this mode. + * Token auth (``IAM_TOKEN_DB_AUTH`` truthy for AWS RDS IAM, or + ``AZURE_POSTGRESQL_AUTH`` truthy for Azure Database for PostgreSQL with + Microsoft Entra ID): mint a short-lived token and embed it as the + password. The writer URL is always (re)written because the token is + freshly minted on every startup. The chart omits ``DATABASE_PASSWORD`` + in this mode. Enabling both toggles is a startup error. * Password auth: build a percent-encoded URL from ``DATABASE_PASSWORD``. The chart emits the discrete ``DATABASE_*`` fields (never a pre-assembled URL), so URL-reserved characters in the password survive @@ -22,27 +24,37 @@ endpoints are covered: one an operator pinned via ``extraEnv`` — is left untouched and wins. The read replica is opt-in via ``DATABASE_HOST_READ_REPLICA`` and never -clobbers a pre-existing ``DATABASE_URL_READ_REPLICA``, so an IAM writer can -run alongside a password-auth reader (or a precomputed reader URL). Reader -IAM is gated on the single global ``IAM_TOKEN_DB_AUTH`` flag — the chart -only emits the reader IAM env vars when the writer also uses IAM auth. +clobbers a pre-existing ``DATABASE_URL_READ_REPLICA``, so a token-auth writer +can run alongside a password-auth reader (or a precomputed reader URL). Reader +token auth is gated on the same global toggle as the writer: the chart only +emits the reader token env vars when the writer also uses token auth. Reader-side fields fall back to the writer's user / name / schema / port / password when their ``*_READ_REPLICA`` counterpart is unset. """ import os import urllib.parse -from typing import Final, cast +from functools import partial +from typing import Annotated, Final, cast -from pydantic import AliasChoices, Field +from pydantic import AliasChoices, BeforeValidator, Field from pydantic_settings import BaseSettings, SettingsConfigDict -# Imported as a module (not `from ... import generate_iam_auth_token`) so the -# AWS-touching token mint stays patchable at its canonical location in tests. -from litellm.proxy.auth import rds_iam_token +from litellm.proxy.db.token_auth import ( + AZURE_POSTGRESQL_AUTH_ENV_VAR, + DEFAULT_POSTGRES_PORT, + IAM_TOKEN_DB_AUTH_ENV_VAR, + DatabaseTokenAuth, + IAMEndpoint, + build_database_token_auth, + mint_database_token, + token_auth_flag_enabled, +) -_IAM_ENV_KEY: Final = "IAM_TOKEN_DB_AUTH" -_DEFAULT_PG_PORT: Final = "5432" +IamTokenAuthFlag = Annotated[bool, BeforeValidator(partial(token_auth_flag_enabled, env_var=IAM_TOKEN_DB_AUTH_ENV_VAR))] +AzureTokenAuthFlag = Annotated[ + bool, BeforeValidator(partial(token_auth_flag_enabled, env_var=AZURE_POSTGRESQL_AUTH_ENV_VAR)) +] # schema.prisma pins `provider = "postgresql"`, so these are the only schemes # Prisma can actually connect with. @@ -90,13 +102,14 @@ class DatabaseURLSettings(BaseSettings): model_config = SettingsConfigDict(case_sensitive=False, extra="ignore") - iam_token_db_auth: bool = Field(default=False, validation_alias=_IAM_ENV_KEY) + iam_token_db_auth: IamTokenAuthFlag = Field(default=False, validation_alias=IAM_TOKEN_DB_AUTH_ENV_VAR) + azure_postgresql_auth: AzureTokenAuthFlag = Field(default=False, validation_alias=AZURE_POSTGRESQL_AUTH_ENV_VAR) # Writer database_url: str | None = Field(default=None, validation_alias="DATABASE_URL") direct_url: str | None = Field(default=None, validation_alias="DIRECT_URL") database_host: str | None = Field(default=None, validation_alias="DATABASE_HOST") - database_port: str = Field(default=_DEFAULT_PG_PORT, validation_alias="DATABASE_PORT") + database_port: str = Field(default=DEFAULT_POSTGRES_PORT, validation_alias="DATABASE_PORT") database_user: str | None = Field( default=None, validation_alias=AliasChoices("DATABASE_USER", "DATABASE_USERNAME"), @@ -122,15 +135,27 @@ class DatabaseURLSettings(BaseSettings): """Load the settings from ``os.environ`` (read at call time).""" return cls() + def token_auth(self) -> DatabaseTokenAuth | None: + """The token strategy the toggles ask for, or ``None`` for password auth. + + Raises ``RuntimeError`` when both toggles are on, since the password can only + come from one source. + """ + return build_database_token_auth( + iam_token_db_auth=self.iam_token_db_auth, + azure_postgresql_auth=self.azure_postgresql_auth, + ) + def build_writer_url(self) -> str | None: """Return the writer URL to set, or ``None`` to leave it as-is. - Raises ``RuntimeError`` (naming the offending vars) when IAM auth is + Raises ``RuntimeError`` (naming the offending vars) when token auth is enabled but a required field is missing — the proxy cannot recover from this and a clear startup error beats a Prisma connect failure. """ - if self.iam_token_db_auth: - missing: Final = [ + auth: Final = self.token_auth() + if auth is not None: + missing: Final = tuple( env for env, val in ( ("DATABASE_HOST", self.database_host), @@ -138,23 +163,21 @@ class DatabaseURLSettings(BaseSettings): ("DATABASE_NAME", self.database_name), ) if not val - ] + ) if missing: raise RuntimeError( - "IAM_TOKEN_DB_AUTH is enabled but required DB env var(s) " + f"{auth.env_var} is enabled but required DB env var(s) " f"are unset: {', '.join(missing)}. Set them so the writer " - "DATABASE_URL can be assembled with a minted IAM token." + f"DATABASE_URL can be assembled with a minted {auth.label}." ) - host: Final = cast(str, self.database_host) - user: Final = cast(str, self.database_user) - name: Final = cast(str, self.database_name) - # IAM token is already URL-quoted by generate_iam_auth_token; - # user/name embedded raw (parity with proxy_cli.py / IAMEndpoint). - token: Final = rds_iam_token.generate_iam_auth_token(db_host=host, db_port=self.database_port, db_user=user) - url = f"postgresql://{user}:{token}@{host}:{self.database_port}/{name}" - if self.database_schema: - url += f"?schema={self.database_schema}" - return url + endpoint: Final = IAMEndpoint( + host=cast(str, self.database_host), + port=self.database_port, + user=cast(str, self.database_user), + name=cast(str, self.database_name), + schema=self.database_schema, + ) + return endpoint.build_url(mint_database_token(auth, endpoint)) # Password auth: an operator-pinned DATABASE_URL always wins. if self.database_url: @@ -184,35 +207,37 @@ class DatabaseURLSettings(BaseSettings): host: Final = self.database_host_read_replica port: Final = self.database_port_read_replica or self.database_port - user = self.database_user_read_replica or self.database_user - name = self.database_name_read_replica or self.database_name + user: Final = self.database_user_read_replica or self.database_user + name: Final = self.database_name_read_replica or self.database_name schema: Final = self.database_schema_read_replica or self.database_schema password: Final = self.database_password_read_replica or self.database_password - if self.iam_token_db_auth: - missing: Final = [ + auth: Final = self.token_auth() + if auth is not None: + missing: Final = tuple( env for env, val in ( ("DATABASE_USER[_READ_REPLICA]", user), ("DATABASE_NAME[_READ_REPLICA]", name), ) if not val - ] + ) if missing: raise RuntimeError( - "IAM_TOKEN_DB_AUTH is enabled and DATABASE_HOST_READ_REPLICA " + f"{auth.env_var} is enabled and DATABASE_HOST_READ_REPLICA " "is set, but the reader could not resolve: " f"{', '.join(missing)} (no *_READ_REPLICA value and no " "writer fallback). Set the reader fields or the writer " "defaults." ) - user = cast(str, user) - name = cast(str, name) - token: Final = rds_iam_token.generate_iam_auth_token(db_host=host, db_port=port, db_user=user) - url = f"postgresql://{user}:{token}@{host}:{port}/{name}" - if schema: - url += f"?schema={schema}" - return url + endpoint: Final = IAMEndpoint( + host=host, + port=port, + user=cast(str, user), + name=cast(str, name), + schema=schema, + ) + return endpoint.build_url(mint_database_token(auth, endpoint)) if user and name: return self._password_url( @@ -271,23 +296,35 @@ class DatabaseURLSettings(BaseSettings): if bad_scheme is not None: raise RuntimeError(unsupported_db_scheme_message(env_var, bad_scheme)) + def apply_writer_url_to_env(self) -> bool: + """Write just the assembled writer URL into ``os.environ``. + + Split out because the CLI shares this minting path but resolves the read + replica separately, so it must not pick up reader behavior on the way. The + CLI runs its own scheme guard over the pinned URLs, so unlike + ``apply_to_env`` this does not repeat it. + """ + writer_url: Final = self.build_writer_url() + if writer_url is None: + return False + os.environ["DATABASE_URL"] = writer_url + # Normalize the toggles so downstream readers (PrismaWrapper's token + # refresh) reliably see token auth on, regardless of spelling. + if self.iam_token_db_auth: + os.environ[IAM_TOKEN_DB_AUTH_ENV_VAR] = "True" + if self.azure_postgresql_auth: + os.environ[AZURE_POSTGRESQL_AUTH_ENV_VAR] = "True" + return True + def apply_to_env(self) -> bool: """Write the assembled URL(s) into ``os.environ``. - Returns True iff this call set ``DATABASE_URL`` (IAM mint, or + Returns True iff this call set ``DATABASE_URL`` (token mint, or password auth that assembled a fresh URL). False means there was nothing to do — an operator-pinned URL, or no discrete fields. """ self._raise_for_unsupported_scheme() - wrote_writer = False - writer_url: Final = self.build_writer_url() - if writer_url is not None: - os.environ["DATABASE_URL"] = writer_url - if self.iam_token_db_auth: - # Normalize the toggle so downstream readers (PrismaWrapper's - # IAM refresh) reliably see IAM on, regardless of spelling. - os.environ[_IAM_ENV_KEY] = "True" - wrote_writer = True + wrote_writer: Final = self.apply_writer_url_to_env() reader_url: Final = self.build_reader_url() if reader_url is not None: diff --git a/litellm/proxy/db/exception_handler.py b/litellm/proxy/db/exception_handler.py index f7a39aaa50f..5502543b926 100644 --- a/litellm/proxy/db/exception_handler.py +++ b/litellm/proxy/db/exception_handler.py @@ -335,6 +335,7 @@ async def call_with_db_reconnect_retry( coro_factory: Callable[[], Awaitable[_ReadResultT]], *, reason: str, + retry_safe_error_types: tuple[type[Exception], ...] | None = None, timeout_seconds: float | None = None, lock_timeout_seconds: float | None = None, ) -> _ReadResultT: @@ -350,7 +351,8 @@ async def call_with_db_reconnect_retry( 2. On exception, if it is NOT a transport error (per `is_database_transport_error`), re-raise — data-layer errors like `UniqueViolationError` mean the DB is reachable, reconnect would be - pointless. + pointless. Transport errors outside `retry_safe_error_types` are + re-raised too. 3. If `prisma_client` does not expose `attempt_db_reconnect`, re-raise. This guards against partial stand-ins / older clients in tests. 4. Call `prisma_client.attempt_db_reconnect(reason=...)`. If it returns @@ -371,6 +373,10 @@ async def call_with_db_reconnect_retry( `attempt_db_reconnect` and the `_db_auth_reconnect_*` defaults. coro_factory: Zero-arg callable returning the read awaitable. reason: Telemetry tag forwarded to `attempt_db_reconnect`. + retry_safe_error_types: Which transport errors may be replayed, or + None for every transport error. A non-idempotent write must narrow + this to `DB_RETRY_SAFE_ERROR_TYPES`, where the statements provably + never reached the database. timeout_seconds: Optional override for the reconnect cycle timeout. Defaults to `prisma_client._db_auth_reconnect_timeout_seconds`, then to 2.0s. @@ -392,6 +398,8 @@ async def call_with_db_reconnect_retry( except Exception as first_exc: if not PrismaDBExceptionHandler.is_database_transport_error(first_exc): raise + if retry_safe_error_types is not None and not isinstance(first_exc, retry_safe_error_types): + raise if not hasattr(prisma_client, "attempt_db_reconnect"): raise diff --git a/litellm/proxy/db/prisma_client.py b/litellm/proxy/db/prisma_client.py index 5f86490a474..fc761fc1831 100644 --- a/litellm/proxy/db/prisma_client.py +++ b/litellm/proxy/db/prisma_client.py @@ -1,5 +1,6 @@ """ -This file contains the PrismaWrapper class, which is used to wrap the Prisma client and handle the RDS IAM token. +This file contains the PrismaWrapper class, which wraps the Prisma client and keeps the +database token (AWS RDS IAM or Microsoft Entra ID) fresh. """ import asyncio @@ -11,34 +12,27 @@ import time import urllib import urllib.parse from collections.abc import Callable -from dataclasses import dataclass from datetime import datetime, timedelta from typing import Any, Final, Protocol from litellm._logging import verbose_proxy_logger +from litellm.proxy.db.token_auth import ( + DEFAULT_POSTGRES_PORT, + DatabaseTokenAuth, + IAMEndpoint, + RdsIamTokenAuth, + mint_database_token, + parse_database_token_expiration, + parse_iam_endpoint_from_url, +) from litellm.secret_managers.main import str_to_bool - -@dataclass(frozen=True) -class IAMEndpoint: - """Static parts of an RDS IAM-authenticated Postgres connection. - - The IAM token rotates every ~15 minutes; everything else (host, port, user, - database name, schema) stays fixed. We capture the static fields once so - refresh just regenerates the token and reassembles the URL. - """ - - host: str - port: str - user: str - name: str - schema: str | None = None - - def build_url(self, token: str) -> str: - url = f"postgresql://{self.user}:{token}@{self.host}:{self.port}/{self.name}" - if self.schema: - url += f"?schema={self.schema}" - return url +__all__ = ( + "IAMEndpoint", + "PrismaManager", + "PrismaWrapper", + "parse_iam_endpoint_from_url", +) class _PrismaProcess(Protocol): @@ -141,45 +135,17 @@ class _TrackedPrismaEngine: self.tracker.transaction_finished(tx_id) -def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint: - """Parse an IAMEndpoint from a Postgres URL. - - Used so a reader URL can drive its own IAM refresh without requiring - callers to set parallel DATABASE_HOST_READ_REPLICA / etc. env vars. - """ - parsed: Final = urllib.parse.urlparse(url) - if not parsed.hostname or not parsed.username: - raise ValueError("Cannot parse IAM endpoint from URL: missing host or username") - name: Final = (parsed.path or "/").lstrip("/") - if not name: - raise ValueError("Cannot parse IAM endpoint from URL: missing database name") - port: Final = str(parsed.port) if parsed.port else "5432" - schema: str | None = None - if parsed.query: - qs: Final = urllib.parse.parse_qs(parsed.query) - schema_vals: Final = qs.get("schema") - if schema_vals: - schema = schema_vals[0] - return IAMEndpoint( - host=parsed.hostname, - port=port, - user=parsed.username, - name=name, - schema=schema, - ) - - class PrismaWrapper: """ - Wrapper around Prisma client that handles RDS IAM token authentication. + Wrapper around Prisma client that handles token-based database authentication. - When iam_token_db_auth is enabled, this wrapper: - 1. Proactively refreshes IAM tokens before they expire (background task) + When a token strategy is active (AWS RDS IAM or Microsoft Entra ID), this wrapper: + 1. Proactively refreshes the token before it expires (background task) 2. Falls back to synchronous refresh if a token is found expired 3. Uses proper locking to prevent race conditions during reconnection - RDS IAM tokens are valid for 15 minutes. This wrapper refreshes them - 3 minutes before expiration to ensure uninterrupted database connectivity. + RDS IAM tokens are valid for 15 minutes and Entra tokens for about an hour. This + wrapper refreshes 3 minutes before whatever expiry the live token carries. """ # Buffer time in seconds before token expiration to trigger refresh @@ -189,20 +155,28 @@ class PrismaWrapper: # Fallback refresh interval if token parsing fails (10 minutes) FALLBACK_REFRESH_INTERVAL_SECONDS = 600 + # Floor on the proactive loop's sleep, so a token whose expiry does not advance + # (azure-identity hands back its cached token when a renewal attempt fails) costs + # one retry every 30 seconds instead of spinning the loop with no sleep at all. + TOKEN_REFRESH_MIN_SLEEP_SECONDS = 30 + ENGINE_RETIREMENT_DRAIN_TIMEOUT_SECONDS = 90 def __init__( self, original_prisma: Any, - iam_token_db_auth: bool, + iam_token_db_auth: bool = False, *, + token_auth: DatabaseTokenAuth | None = None, db_url_env_var: str = "DATABASE_URL", iam_endpoint: IAMEndpoint | None = None, recreate_uses_datasource: bool = False, log_prefix: str = "", ): + # Set before `_original_prisma` so the `iam_token_db_auth` property below can + # never send `__getattr__` looking for a half-built strategy on the raw client. + self._token_auth = token_auth if token_auth is not None else (RdsIamTokenAuth() if iam_token_db_auth else None) self._original_prisma = original_prisma - self.iam_token_db_auth = iam_token_db_auth # Per-connection knobs so the same wrapper can be used for the writer # (defaults: DATABASE_URL env, IAM endpoint from DATABASE_HOST/etc., @@ -241,6 +215,25 @@ class PrismaWrapper: self._engine_generation: int = 0 self.on_engine_replaced: Callable[[], None] | None = None + @property + def token_auth(self) -> DatabaseTokenAuth | None: + """The active database token strategy, or None for password auth.""" + return self._token_auth + + @property + def token_label(self) -> str: + """Human name of the active token kind, for log lines.""" + return self._token_auth.label if self._token_auth is not None else "database token" + + @property + def iam_token_db_auth(self) -> bool: + """Whether any token strategy is active. + + Read-only: the kind of token is chosen once, by injection, so there is no way + to flip this back on and silently get AWS RDS on an Azure deployment. + """ + return self._token_auth is not None + @staticmethod def _read_engine(prisma_client: _PrismaClient) -> _PrismaEngine: return prisma_client._engine @@ -376,30 +369,9 @@ class PrismaWrapper: Returns the datetime when the token expires, or None if parsing fails. """ - if token is None: - return None - - try: - # Token format: ...?X-Amz-Date=YYYYMMDDTHHMMSSZ&X-Amz-Expires=900&... - if "?" not in token: - return None - - query_string: Final = token.split("?", 1)[1] - params: Final = urllib.parse.parse_qs(query_string) - - expires_str: Final = params.get("X-Amz-Expires", [None])[0] - date_str: Final = params.get("X-Amz-Date", [None])[0] - - if not expires_str or not date_str: - return None - - token_created: Final = datetime.strptime(date_str, "%Y%m%dT%H%M%SZ") - expires_in: Final = int(expires_str) - - return token_created + timedelta(seconds=expires_in) - except Exception as e: - verbose_proxy_logger.debug("Failed to parse token expiration: %s", e) + if token is None or self._token_auth is None: return None + return parse_database_token_expiration(self._token_auth, token) def _calculate_seconds_until_refresh(self) -> float: """ @@ -409,8 +381,9 @@ class PrismaWrapper: For a 15-minute (900s) token with 180s buffer, this returns ~720s (12 min). Returns: - Number of seconds to sleep before the next refresh. - Returns 0 if token should be refreshed immediately. + Number of seconds to sleep before the next refresh, never less than + TOKEN_REFRESH_MIN_SLEEP_SECONDS so a token whose expiry never advances + cannot spin the loop. Returns FALLBACK_REFRESH_INTERVAL_SECONDS if parsing fails. """ db_url: Final = os.getenv(self._db_url_env_var) @@ -432,8 +405,10 @@ class PrismaWrapper: now: Final = datetime.utcnow() seconds_until_refresh: Final = (refresh_at - now).total_seconds() - # If already past refresh time, return 0 (refresh immediately) - return max(0, seconds_until_refresh) + # Past refresh time means refresh as soon as the floor allows, not instantly: + # a provider that keeps handing back the same token would otherwise leave the + # loop re-minting and recreating the query engine with no sleep between passes. + return max(self.TOKEN_REFRESH_MIN_SLEEP_SECONDS, seconds_until_refresh) def is_token_expired(self, token_url: str | None) -> bool: """Check if the token in the given URL is expired.""" @@ -451,40 +426,47 @@ class PrismaWrapper: return datetime.utcnow() > expiration_time def get_rds_iam_token(self) -> str | None: - """Generate a new RDS IAM token and update the configured DB URL env var. + """Mint a fresh database token and update the configured DB URL env var. When the wrapper was constructed with an explicit `iam_endpoint` (typical for a reader wrapper whose host/port/user came from a parsed - URL), use that. Otherwise fall back to the legacy DATABASE_HOST/PORT/ - USER/NAME/SCHEMA env vars (writer behavior). + URL), use that. Otherwise fall back to the DATABASE_HOST/PORT/USER/ + NAME/SCHEMA env vars (writer behavior). """ - if not self.iam_token_db_auth: + auth: Final = self._token_auth + if auth is None: return None - from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token + endpoint: Final = self._iam_endpoint if self._iam_endpoint is not None else self._endpoint_from_env() + db_url: Final = endpoint.build_url(mint_database_token(auth, endpoint)) + os.environ[self._db_url_env_var] = db_url + return db_url - if self._iam_endpoint is not None: - endpoint: Final = self._iam_endpoint - token = generate_iam_auth_token(db_host=endpoint.host, db_port=endpoint.port, db_user=endpoint.user) - _db_url = endpoint.build_url(token) - else: - db_host: Final = os.getenv("DATABASE_HOST") + @staticmethod + def _endpoint_from_env() -> IAMEndpoint: + host: Final = os.getenv("DATABASE_HOST") + user: Final = os.getenv("DATABASE_USER") + name: Final = os.getenv("DATABASE_NAME") + if not host or not user or not name: + missing: Final = tuple( + env + for env, value in (("DATABASE_HOST", host), ("DATABASE_USER", user), ("DATABASE_NAME", name)) + if not value + ) + raise RuntimeError( + f"Cannot mint a database token: {', '.join(missing)} unset. Set them so the " + "connection URL can be reassembled around a freshly minted token." + ) + return IAMEndpoint( + host=host, # Default to the Postgres standard port; passing None to # `generate_iam_auth_token` makes botocore embed the literal # string "None" in the presigned URL, which then fails to parse. - db_port: Final = os.getenv("DATABASE_PORT", "5432") - db_user: Final = os.getenv("DATABASE_USER") - db_name: Final = os.getenv("DATABASE_NAME") - db_schema: Final = os.getenv("DATABASE_SCHEMA") - - token = generate_iam_auth_token(db_host=db_host, db_port=db_port, db_user=db_user) - - _db_url = f"postgresql://{db_user}:{token}@{db_host}:{db_port}/{db_name}" - if db_schema: - _db_url += f"?schema={db_schema}" - - os.environ[self._db_url_env_var] = _db_url - return _db_url + port=os.getenv("DATABASE_PORT", DEFAULT_POSTGRES_PORT), + user=user, + name=name, + schema=os.getenv("DATABASE_SCHEMA"), + ) @property def engine_generation(self) -> int: @@ -658,12 +640,12 @@ class PrismaWrapper: """ Start the background token refresh task. - This task proactively refreshes RDS IAM tokens before they expire, + This task proactively refreshes the database token before it expires, preventing connection failures. Should be called after the initial Prisma client connection is established. """ if not self.iam_token_db_auth: - verbose_proxy_logger.debug("IAM token auth not enabled, skipping token refresh task") + verbose_proxy_logger.debug("Database token auth not enabled, skipping token refresh task") return if self._token_refresh_task is not None: @@ -672,8 +654,9 @@ class PrismaWrapper: self._token_refresh_task = asyncio.create_task(self._token_refresh_loop()) verbose_proxy_logger.info( - "%sStarted RDS IAM token proactive refresh background task", + "%sStarted %s proactive refresh background task", self._log_prefix, + self.token_label, ) async def stop_token_refresh_task(self) -> None: @@ -691,19 +674,24 @@ class PrismaWrapper: except asyncio.CancelledError: pass self._token_refresh_task = None - verbose_proxy_logger.info("%sStopped RDS IAM token refresh background task", self._log_prefix) + verbose_proxy_logger.info( + "%sStopped %s refresh background task", + self._log_prefix, + self.token_label, + ) async def _token_refresh_loop(self) -> None: """ - Background loop that proactively refreshes RDS IAM tokens before expiration. + Background loop that proactively refreshes database tokens before expiration. Uses precise timing: calculates the exact sleep duration until the token needs to be refreshed (expiration - 3 minute buffer), then refreshes. This is more efficient than polling, requiring only 1 wake-up per token cycle. """ verbose_proxy_logger.info( - "%sRDS IAM token refresh loop started. Tokens will be refreshed %ss before expiration.", + "%s%s refresh loop started. Tokens will be refreshed %ss before expiration.", self._log_prefix, + self.token_label, self.TOKEN_REFRESH_BUFFER_SECONDS, ) @@ -714,22 +702,31 @@ class PrismaWrapper: if sleep_seconds > 0: verbose_proxy_logger.info( - f"{self._log_prefix}RDS IAM token refresh scheduled in " + f"{self._log_prefix}{self.token_label} refresh scheduled in " f"{sleep_seconds:.0f} seconds ({sleep_seconds / 60:.1f} minutes)" ) await asyncio.sleep(sleep_seconds) # Refresh the token - verbose_proxy_logger.info("%sProactively refreshing RDS IAM token...", self._log_prefix) + verbose_proxy_logger.info( + "%sProactively refreshing %s...", + self._log_prefix, + self.token_label, + ) await self._safe_refresh_token() except asyncio.CancelledError: - verbose_proxy_logger.info("%sRDS IAM token refresh loop cancelled", self._log_prefix) + verbose_proxy_logger.info( + "%s%s refresh loop cancelled", + self._log_prefix, + self.token_label, + ) break except Exception as e: verbose_proxy_logger.error( - "%sError in RDS IAM token refresh loop: %s. Retrying in %ss...", + "%sError in %s refresh loop: %s. Retrying in %ss...", self._log_prefix, + self.token_label, e, self.FALLBACK_REFRESH_INTERVAL_SECONDS, ) @@ -741,7 +738,7 @@ class PrismaWrapper: async def _safe_refresh_token(self) -> None: """ - Refresh the RDS IAM token with proper locking to prevent race conditions. + Refresh the database token with proper locking to prevent race conditions. Uses an asyncio lock to ensure only one refresh operation happens at a time, preventing multiple concurrent reconnection attempts. @@ -754,8 +751,9 @@ class PrismaWrapper: # by skipping when the current token still has comfortable runway. if self._token_refresh_not_needed(os.getenv(self._db_url_env_var)): verbose_proxy_logger.debug( - "%sRDS IAM token still fresh; skipping redundant refresh.", + "%s%s still fresh; skipping redundant refresh.", self._log_prefix, + self.token_label, ) return @@ -772,13 +770,15 @@ class PrismaWrapper: raise self._last_refresh_time = datetime.utcnow() verbose_proxy_logger.info( - "%sRDS IAM token refreshed successfully. New token valid for ~15 minutes.", + "%s%s refreshed successfully.", self._log_prefix, + self.token_label, ) else: verbose_proxy_logger.error( - "%sFailed to generate new RDS IAM token during proactive refresh", + "%sFailed to generate new %s during proactive refresh", self._log_prefix, + self.token_label, ) def _token_refresh_not_needed(self, token_url: str | None) -> bool: @@ -832,10 +832,11 @@ class PrismaWrapper: if running_loop is not None: verbose_proxy_logger.warning( - "%sRDS IAM token expired in __getattr__ — proactive refresh " + "%s%s expired in __getattr__ - proactive refresh " "may have failed. Scheduling async refresh; the current " "request may fail and be retried with the fresh token.", self._log_prefix, + self.token_label, ) # Non-blocking: schedule the locked refresh on the # running loop. The reconnection lock inside @@ -843,9 +844,10 @@ class PrismaWrapper: running_loop.create_task(self._safe_refresh_token()) else: verbose_proxy_logger.warning( - "%sRDS IAM token expired in __getattr__ — proactive refresh " + "%s%s expired in __getattr__ - proactive refresh " "may have failed. Triggering synchronous fallback refresh...", self._log_prefix, + self.token_label, ) new_db_url: Final = self.get_rds_iam_token() if new_db_url: @@ -857,7 +859,7 @@ class PrismaWrapper: self._log_prefix, ) else: - raise ValueError("Failed to get RDS IAM token") + raise ValueError(f"Failed to get {self.token_label}") return original_attr diff --git a/litellm/proxy/db/routing_prisma_wrapper.py b/litellm/proxy/db/routing_prisma_wrapper.py index 5aeb52be535..22fc32a898a 100644 --- a/litellm/proxy/db/routing_prisma_wrapper.py +++ b/litellm/proxy/db/routing_prisma_wrapper.py @@ -248,14 +248,14 @@ class RoutingPrismaWrapper: async def _recreate_reader(self, http_client: Any | None = None) -> None: """Resolve the reader URL and recreate its Prisma client. - IAM-enabled readers regenerate their token (host/port/user came from - the parsed reader URL at construction time). Non-IAM readers reuse - the URL stored in `DATABASE_URL_READ_REPLICA`. + Token-authenticated readers regenerate their token (host/port/user came + from the parsed reader URL at construction time). Password-authenticated + readers reuse the URL stored in `DATABASE_URL_READ_REPLICA`. """ if self._reader.iam_token_db_auth: new_reader_url: Final = self._reader.get_rds_iam_token() if not new_reader_url: - raise RuntimeError("Failed to generate fresh IAM token for read replica") + raise RuntimeError(f"Failed to generate fresh {self._reader.token_label} for read replica") await self._reader.recreate_prisma_client(new_reader_url, http_client=http_client) return reader_url: Final = os.getenv("DATABASE_URL_READ_REPLICA", "") diff --git a/litellm/proxy/db/token_auth.py b/litellm/proxy/db/token_auth.py new file mode 100644 index 00000000000..e1f84d1c04c --- /dev/null +++ b/litellm/proxy/db/token_auth.py @@ -0,0 +1,274 @@ +"""Token-based authentication for the proxy's Postgres connection. + +Two managed Postgres offerings hand the client a short-lived credential that is used as +the Postgres password: AWS RDS with IAM auth, and Azure Database for PostgreSQL Flexible +Server with Microsoft Entra ID. Both need the same machinery (mint at startup, read the +expiry back off the token, mint again before it lapses) and differ only in how the token +is produced and how its expiry is encoded, so the difference lives in a tagged union that +is resolved once from the environment and injected into whatever needs a token. +""" + +import base64 +import functools +import os +import urllib.parse +from collections.abc import Callable +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import Final, TypeAlias + +from pydantic import BaseModel +from typing_extensions import assert_never + +from litellm._logging import verbose_proxy_logger + +IAM_TOKEN_DB_AUTH_ENV_VAR: Final = "IAM_TOKEN_DB_AUTH" +AZURE_POSTGRESQL_AUTH_ENV_VAR: Final = "AZURE_POSTGRESQL_AUTH" +AZURE_POSTGRESQL_SCOPE: Final = "https://ossrdbms-aad.database.windows.net/.default" + +CONFLICTING_TOKEN_AUTH_MESSAGE: Final = ( + f"{IAM_TOKEN_DB_AUTH_ENV_VAR} and {AZURE_POSTGRESQL_AUTH_ENV_VAR} are both enabled, but the " + "database password can only come from one token source. Keep " + f"{IAM_TOKEN_DB_AUTH_ENV_VAR} for AWS RDS IAM auth, or {AZURE_POSTGRESQL_AUTH_ENV_VAR} for " + "Azure Database for PostgreSQL with Microsoft Entra ID, and unset the other one." +) + +DEFAULT_POSTGRES_PORT: Final = "5432" + +TRUTHY_TOKEN_AUTH_VALUES: Final[frozenset[str]] = frozenset({"1", "on", "t", "true", "y", "yes"}) +FALSY_TOKEN_AUTH_VALUES: Final[frozenset[str]] = frozenset({"", "0", "f", "false", "n", "no", "off"}) + + +def token_auth_flag_enabled(value: str | bool | None, *, env_var: str) -> bool: + """Whether a token-auth toggle is on, rejecting anything it cannot read. + + The single parser for both toggles. Every entry point (the settings model, the + CLI, and the refresh loop's own env lookup) routes through this, so a value like + ``"1"`` cannot enable minting in one place and leave the refresh loop convinced + token auth is off, which would strand a pod on a token it never renews. + + A value that is neither recognizably on nor recognizably off raises: silently + reading a typo as off would downgrade an operator from token auth to password + auth, and the first sign of it would be a connection refused by the server. + """ + if isinstance(value, bool): + return value + if value is None: + return False + normalized: Final = value.strip().lower() + if normalized in TRUTHY_TOKEN_AUTH_VALUES: + return True + if normalized in FALSY_TOKEN_AUTH_VALUES: + return False + raise ValueError( + f"{env_var}={value!r} is not a recognized boolean. Set it to one of " + f"{', '.join(sorted(TRUTHY_TOKEN_AUTH_VALUES))} to turn token auth on, or to one of " + f"{', '.join(sorted(v for v in FALSY_TOKEN_AUTH_VALUES if v))} to turn it off." + ) + + +def _quote(value: str) -> str: + return urllib.parse.quote(value, safe="") + + +def _normalize_quote(value: str) -> str: + """Percent-encode a URL component that may already be percent-encoded. + + ``DATABASE_USER`` used to be interpolated raw, so pre-encoding was the only way to + put an ``@`` in it. Encoding such a value again would double-escape it, so decode + first: the round trip is idempotent and leaves an already-encoded value byte for + byte as it was, while a raw UPN like ``svc@corp`` still comes out encoded. + """ + return urllib.parse.quote(urllib.parse.unquote(value), safe="") + + +@dataclass(frozen=True, slots=True) +class IAMEndpoint: + """Static parts of a token-authenticated Postgres connection. + + The token rotates every few minutes to an hour depending on the provider; + everything else (host, port, user, database name, schema) stays fixed. Capturing + the static fields once means a refresh only regenerates the token and reassembles + the URL. + """ + + host: str + port: str + user: str + name: str + schema: str | None = None + + def build_url(self, token: str) -> str: + """Assemble the connection URL, inserting ``token`` verbatim as the password. + + User, database name, and schema are normalized rather than encoded outright, + because an Entra principal is a UPN containing ``@`` while an operator on the + older RDS path may already have encoded that ``@`` themselves. The token is + left alone: both providers hand it back already in wire form, and re-encoding + it would double-escape the password. + """ + base: Final = ( + f"postgresql://{_normalize_quote(self.user)}:{token}@{self.host}:{self.port}/{_normalize_quote(self.name)}" + ) + if not self.schema: + return base + return f"{base}?schema={_normalize_quote(self.schema)}" + + +def parse_iam_endpoint_from_url(url: str) -> IAMEndpoint: + """Parse an :class:`IAMEndpoint` back out of a Postgres URL. + + Used so a reader URL can drive its own token refresh without requiring callers to + set parallel ``DATABASE_HOST_READ_REPLICA`` / etc. env vars. + """ + parsed: Final = urllib.parse.urlparse(url) + if not parsed.hostname or not parsed.username: + raise ValueError("Cannot parse IAM endpoint from URL: missing host or username") + name: Final = urllib.parse.unquote((parsed.path or "/").lstrip("/")) + if not name: + raise ValueError("Cannot parse IAM endpoint from URL: missing database name") + port: Final = str(parsed.port) if parsed.port else DEFAULT_POSTGRES_PORT + schema_values: Final = urllib.parse.parse_qs(parsed.query).get("schema") if parsed.query else None + return IAMEndpoint( + host=parsed.hostname, + port=port, + user=urllib.parse.unquote(parsed.username), + name=name, + schema=schema_values[0] if schema_values else None, + ) + + +@dataclass(frozen=True, slots=True) +class RdsIamTokenAuth: + """AWS RDS IAM auth: a SigV4-presigned token minted from the ambient AWS credentials.""" + + @property + def label(self) -> str: + return "RDS IAM token" + + @property + def env_var(self) -> str: + return IAM_TOKEN_DB_AUTH_ENV_VAR + + +@dataclass(frozen=True, slots=True) +class AzureEntraTokenAuth: + """Azure Database for PostgreSQL auth: a Microsoft Entra ID access token as the password. + + The provider is injected rather than resolved here so callers (and tests) decide which + Azure credential mints the token. + """ + + token_provider: Callable[[], str] + + @property + def label(self) -> str: + return "Azure Entra token" + + @property + def env_var(self) -> str: + return AZURE_POSTGRESQL_AUTH_ENV_VAR + + +DatabaseTokenAuth: TypeAlias = RdsIamTokenAuth | AzureEntraTokenAuth + + +def mint_database_token(auth: DatabaseTokenAuth, endpoint: IAMEndpoint) -> str: + """Mint a fresh database password for ``endpoint``, already percent-encoded.""" + match auth: + case RdsIamTokenAuth(): + from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token + + return generate_iam_auth_token(db_host=endpoint.host, db_port=endpoint.port, db_user=endpoint.user) + case AzureEntraTokenAuth(): + return _quote(auth.token_provider()) + case _: + assert_never(auth) + + +def parse_database_token_expiration(auth: DatabaseTokenAuth, token: str) -> datetime | None: + """Return when ``token`` expires as a naive UTC datetime, or None when unreadable. + + Callers fall back to a fixed refresh interval on None, so an unparseable token + degrades to periodic refresh instead of failing. + """ + match auth: + case RdsIamTokenAuth(): + return _parse_rds_token_expiration(token) + case AzureEntraTokenAuth(): + return _parse_entra_token_expiration(token) + case _: + assert_never(auth) + + +def _parse_rds_token_expiration(token: str) -> datetime | None: + if "?" not in token: + return None + try: + params: Final = urllib.parse.parse_qs(token.split("?", 1)[1]) + expires_values: Final = params.get("X-Amz-Expires") + date_values: Final = params.get("X-Amz-Date") + if not expires_values or not date_values: + return None + created: Final = datetime.strptime(date_values[0], "%Y%m%dT%H%M%SZ") + return created + timedelta(seconds=int(expires_values[0])) + except (ValueError, OverflowError, OSError) as exc: + verbose_proxy_logger.debug("Failed to parse RDS IAM token expiration: %s", exc) + return None + + +class _EntraAccessTokenClaims(BaseModel): + exp: int + + +def _parse_entra_token_expiration(token: str) -> datetime | None: + segments: Final = token.split(".") + if len(segments) != 3: + return None + payload: Final = segments[1] + try: + claims: Final = _EntraAccessTokenClaims.model_validate_json( + base64.urlsafe_b64decode(payload + "=" * (-len(payload) % 4)) + ) + except ValueError as exc: + verbose_proxy_logger.debug("Failed to parse Azure Entra token expiration: %s", exc) + return None + return datetime.fromtimestamp(claims.exp, tz=timezone.utc).replace(tzinfo=None) + + +@functools.cache +def build_azure_entra_token_provider() -> Callable[[], str]: + """The process-wide Entra token provider for the Azure Postgres OSS RDBMS scope. + + Cached because the writer URL, the reader URL, and the refresh loop each ask for a + strategy, and every uncached call would build another Azure credential with its own + HTTP transport and its own token cache that nothing ever closes. + """ + from litellm.secret_managers.get_azure_ad_token_provider import ( + get_azure_ad_token_provider, + ) + + return get_azure_ad_token_provider(azure_scope=AZURE_POSTGRESQL_SCOPE) + + +def build_database_token_auth(*, iam_token_db_auth: bool, azure_postgresql_auth: bool) -> DatabaseTokenAuth | None: + """Pick the token strategy the two toggles ask for, or None when neither is on.""" + if iam_token_db_auth and azure_postgresql_auth: + raise RuntimeError(CONFLICTING_TOKEN_AUTH_MESSAGE) + if azure_postgresql_auth: + return AzureEntraTokenAuth(token_provider=build_azure_entra_token_provider()) + if iam_token_db_auth: + return RdsIamTokenAuth() + return None + + +def resolve_database_token_auth() -> DatabaseTokenAuth | None: + """Resolve the token strategy from the environment, raising when both toggles are set.""" + return build_database_token_auth( + iam_token_db_auth=token_auth_flag_enabled( + os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR), env_var=IAM_TOKEN_DB_AUTH_ENV_VAR + ), + azure_postgresql_auth=token_auth_flag_enabled( + os.getenv(AZURE_POSTGRESQL_AUTH_ENV_VAR), env_var=AZURE_POSTGRESQL_AUTH_ENV_VAR + ), + ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py index 5bbb01c6c8e..e95e97bfe74 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py +++ b/litellm/proxy/guardrails/guardrail_hooks/unified_guardrail/unified_guardrail.py @@ -38,7 +38,7 @@ if TYPE_CHECKING: BaseTranslation, ) -# Call types that use NDJSON streaming (A2A); guardrail HTTPException is emitted as in-stream error +# Call types that stream JSON-RPC events (A2A); guardrail HTTPException is emitted as in-stream error A2A_CALL_TYPES: Final = (CallTypes.asend_message, CallTypes.send_message) GUARDRAIL_NAME: Final = "unified_llm_guardrails" @@ -90,6 +90,24 @@ def _get_a2a_request_id(responses_so_far: Sequence[object], request_data: dict) return None +def _a2a_jsonrpc_error_chunk(exc: HTTPException, request_id: str | None) -> Mapping[str, object]: + """Build the in-stream JSON-RPC error object for a mid-stream A2A failure. + + Returned as an object, not a serialized string: the A2A endpoint owns wire + framing and serializes whatever the stream yields. + """ + detail: Final = exc.detail if isinstance(exc.detail, dict) else {"message": str(exc.detail)} + return { + "jsonrpc": "2.0", + "id": request_id, + "error": { + "code": -32603, + "message": detail.get("error", detail.get("message", str(exc.detail))), + "data": {k: v for k, v in detail.items() if k not in ("error", "message")}, + }, + } + + endpoint_guardrail_translation_mappings = None @@ -391,28 +409,12 @@ class UnifiedLLMGuardrails(CustomLogger): responses_so_far: Sequence[object], request_data: dict, ) -> AsyncGenerator[object, None]: - """Surface a mid-stream HTTPException. For A2A (NDJSON) call types the - response has already started, so emit an in-stream JSON-RPC error chunk; - otherwise re-raise so the proxy can report it. + """Surface a mid-stream HTTPException. For A2A call types the response has + already started, so emit an in-stream JSON-RPC error chunk; otherwise + re-raise so the proxy can report it. """ if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES: - request_id: Final = _get_a2a_request_id(responses_so_far, request_data) - detail: Final = exc.detail if isinstance(exc.detail, dict) else {"message": str(exc.detail)} - error_chunk: Final = ( - json.dumps( - { - "jsonrpc": "2.0", - "id": request_id, - "error": { - "code": -32603, - "message": detail.get("error", detail.get("message", str(exc.detail))), - "data": {k: v for k, v in detail.items() if k not in ("error", "message")}, - }, - } - ) - + "\n" - ) - yield error_chunk + yield _a2a_jsonrpc_error_chunk(exc, _get_a2a_request_id(responses_so_far, request_data)) return raise exc @@ -1068,28 +1070,9 @@ class UnifiedLLMGuardrails(CustomLogger): return except HTTPException as e: # Response already started (we already yielded chunks); cannot send 400. - # For A2A (NDJSON), yield an in-stream JSON-RPC error so the client sees it. + # For A2A, yield an in-stream JSON-RPC error so the client sees it. if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES: - request_id = _get_a2a_request_id(responses_so_far, request_data) - detail = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)} - error_chunk = ( - json.dumps( - { - "jsonrpc": "2.0", - "id": request_id, - "error": { - "code": -32603, - "message": detail.get( - "error", - detail.get("message", str(e.detail)), - ), - "data": {k: v for k, v in detail.items() if k not in ("error", "message")}, - }, - } - ) - + "\n" - ) - yield error_chunk + yield _a2a_jsonrpc_error_chunk(e, _get_a2a_request_id(responses_so_far, request_data)) return raise chunks_yielded = True @@ -1151,22 +1134,6 @@ class UnifiedLLMGuardrails(CustomLogger): return except HTTPException as e: if call_type is not None and CallTypes(call_type) in A2A_CALL_TYPES: - request_id = _get_a2a_request_id(responses_so_far, request_data) - detail = e.detail if isinstance(e.detail, dict) else {"message": str(e.detail)} - error_chunk = ( - json.dumps( - { - "jsonrpc": "2.0", - "id": request_id, - "error": { - "code": -32603, - "message": detail.get("error", detail.get("message", str(e.detail))), - "data": {k: v for k, v in detail.items() if k not in ("error", "message")}, - }, - } - ) - + "\n" - ) - yield error_chunk + yield _a2a_jsonrpc_error_chunk(e, _get_a2a_request_id(responses_so_far, request_data)) else: raise diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index d47bd7fa311..46aac82473c 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -37,6 +37,7 @@ from litellm.repositories.base_repository import SupportsModelDump from litellm.repositories.team_repository import TeamRepository from litellm.router_strategy.complexity_router import ComplexityRouter from litellm.types.management_endpoints.auto_router_endpoints import ( + SHADOW_EVAL_TURN_VALVE, AutoRouterBenchmarkGroup, AutoRouterBenchmarksResponse, AutoRouterBenchmarkTotals, @@ -662,12 +663,19 @@ _ATTEMPT_AGG_BY_TIER_SQL: Final = "SELECT COALESCE(tier, 'UNCLASSIFIED') AS grp, _ATTEMPT_AGG_BY_MODEL_SQL: Final = "SELECT COALESCE(real_model, 'unknown') AS grp," + _ATTEMPT_AGG_SELECT _ATTEMPT_AGG_BY_LEG_SQL: Final = "SELECT job_id AS grp," + _ATTEMPT_AGG_SELECT +# These guards derive spend from attempt rows, the cross-pod authority; the sampler also +# reads the live counter, so admission can stop before a row-based guard would fire (safe +# direction, and mid-deploy rows from old pods price as judge-only until the deploy ends). _SWEEP_FINISHED_JOBS_SQL: Final = """ UPDATE "LiteLLM_ShadowEvalJob" j SET stopped_at = (NOW() AT TIME ZONE 'utc') WHERE j.api_key_id = ANY($1::text[]) AND j.stopped_at IS NULL AND ( j.ends_at <= (NOW() AT TIME ZONE 'utc') OR (SELECT COUNT(*) FROM "LiteLLM_ShadowEvalAttempt" a WHERE a.job_id = j.id) >= j.max_turns + OR ( + j.max_budget IS NOT NULL + AND (SELECT COALESCE(SUM(a.judge_cost + a.shadow_cost), 0) FROM "LiteLLM_ShadowEvalAttempt" a WHERE a.job_id = j.id) >= j.max_budget + ) ) """ @@ -681,7 +689,7 @@ WHERE job_id = ANY($1::text[]) """ _ATTEMPT_COUNTS_SQL: Final = """ -SELECT a.job_id, COUNT(*)::int AS attempt_count +SELECT a.job_id, COUNT(*)::int AS attempt_count, COALESCE(SUM(a.judge_cost + a.shadow_cost), 0)::float AS spend FROM "LiteLLM_ShadowEvalAttempt" a JOIN "LiteLLM_ShadowEvalJob" j ON j.id = a.job_id WHERE a.job_id = ANY($1::text[]) AND (j.stopped_at IS NULL OR a.created_at <= j.stopped_at) @@ -697,6 +705,10 @@ WHERE group_id = $1 AND stopped_by IS NULL SELECT 1 FROM "LiteLLM_ShadowEvalJob" k WHERE k.group_id = $1 AND k.stopped_at IS NULL AND (SELECT COUNT(*) FROM "LiteLLM_ShadowEvalAttempt" a WHERE a.job_id = k.id) < k.max_turns + AND ( + k.max_budget IS NULL + OR (SELECT COALESCE(SUM(a.judge_cost + a.shadow_cost), 0) FROM "LiteLLM_ShadowEvalAttempt" a WHERE a.job_id = k.id) < k.max_budget + ) ) """ @@ -704,6 +716,7 @@ WHERE group_id = $1 AND stopped_by IS NULL class _AttemptCountRow(BaseModel): job_id: str attempt_count: int + spend: float _ATTEMPT_COUNT_ROWS: Final = TypeAdapter(list[_AttemptCountRow]) @@ -770,6 +783,7 @@ class _LegRow(BaseModel): judge_model: str shadow_percentage: float max_turns: int + max_budget: float | None = None created_at: datetime ends_at: datetime stopped_at: datetime | None = None @@ -789,22 +803,25 @@ class _LegRow(BaseModel): _LEG_ROWS: Final = TypeAdapter(list[_LegRow]) -async def _leg_attempt_counts(prisma_client: "PrismaClient", legs: Sequence[_LegRow]) -> Mapping[str, int]: - """Each leg's attempt count by leg id, judged and errored alike, in one grouped read. - It is the same count the sampler budgets against max_turns, so the derived status - flips to completed exactly when sampling actually ends. A stamped leg's count freezes - at its stopped_at: in-flight attempts that land after the stamp are excluded, so they - can never reclassify a leg that was stopped under budget as budget-spent.""" +async def _leg_attempt_counts(prisma_client: "PrismaClient", legs: Sequence[_LegRow]) -> Mapping[str, _AttemptCountRow]: + """Each leg's attempt count and recorded spend by leg id, judged and errored alike, in + one grouped read. They are the same figures the sampler budgets against max_turns and + max_budget, so the derived status flips to completed exactly when sampling actually + ends. A stamped leg's figures freeze at its stopped_at: in-flight attempts that land + after the stamp are excluded, so they can never reclassify a leg that was stopped + under budget as budget-spent.""" if not legs: return MappingProxyType({}) rows: Final = _ATTEMPT_COUNT_ROWS.validate_python( await _query_raw(prisma_client, _ATTEMPT_COUNTS_SQL, [leg.id for leg in legs]) # mutable-ok: query param or () ) - return MappingProxyType({row.job_id: row.attempt_count for row in rows}) + return MappingProxyType({row.job_id: row for row in rows}) -def _group_response(group_id: str, legs: Sequence[_LegRow], attempt_counts: Mapping[str, int]) -> ShadowEvalJobResponse: +def _group_response( + group_id: str, legs: Sequence[_LegRow], attempt_counts: Mapping[str, _AttemptCountRow] +) -> ShadowEvalJobResponse: """The one constructor of a job response: the caller names the group and passes that group's legs. Config is read off the first leg because every leg carries the same copy, written by one create_many. No caller may serialize a raw row (that would leak a leg id @@ -816,8 +833,10 @@ def _group_response(group_id: str, legs: Sequence[_LegRow], attempt_counts: Mapp ShadowEvalJobKeyResponse( api_key_id=leg.api_key_id, max_turns=leg.max_turns, + max_budget=leg.max_budget, stopped_at=leg.stopped_at, - attempt_count=attempt_counts.get(leg.id, 0), + attempt_count=stats.attempt_count if (stats := attempt_counts.get(leg.id)) else 0, + spend=round(stats.spend, 6) if stats else 0.0, ) for leg in sorted(legs, key=lambda leg: leg.api_key_id) ), @@ -923,11 +942,12 @@ async def start_shadow_eval( serve and duplicates them against baseline_model. A key can hold one active job per direction, so both questions can run at once. - Shadow responses are never served to users. Each key samples until it has judged - max_turns turns of its own traffic, the job's window ends, or the job is stopped, so one - key running out of budget does not end sampling for the others; sampling changes - propagate to pods within about 10 seconds. Shadow and judge calls bill to the shadowed - key but are excluded from request counts and auto-router adoption metrics. + Shadow responses are never served to users. Each key samples until its recorded eval + spend, the shadow and judge calls' own cost, reaches max_budget dollars, the job's + window ends, or the job is stopped, so one key running out of budget does not end + sampling for the others; sampling changes propagate to pods within about 10 seconds. + Shadow and judge calls bill to the shadowed key but are excluded from request counts + and auto-router adoption metrics. """ from litellm.proxy.proxy_server import llm_router, prisma_client @@ -952,7 +972,7 @@ async def start_shadow_eval( ), ) - # A job whose window passed or whose turn budget ran out stopped sampling on its own, + # A job whose window passed or whose budget ran out stopped sampling on its own, # but its legs still hold their slots in the per-key, per-direction partial unique index # until stamped; free them so a new eval can start. Sweeping both directions is deliberate. requested: Final = list(data.api_key_ids) # mutable-ok: query param @@ -983,7 +1003,8 @@ async def start_shadow_eval( "baseline_model": data.baseline_model, "judge_model": data.judge_model, "shadow_percentage": data.shadow_percentage, - "max_turns": data.max_turns, + "max_turns": SHADOW_EVAL_TURN_VALVE, + "max_budget": data.max_budget, "created_by": user_api_key_dict.user_id, "created_at": now, "ends_at": ends_at, @@ -1007,7 +1028,8 @@ async def start_shadow_eval( keys=tuple( ShadowEvalJobKeyResponse( api_key_id=api_key_id, - max_turns=data.max_turns, + max_turns=SHADOW_EVAL_TURN_VALVE, + max_budget=data.max_budget, key_alias=labels[api_key_id].key_alias, key_name=labels[api_key_id].key_name, ) diff --git a/litellm/proxy/management_endpoints/model_management_endpoints.py b/litellm/proxy/management_endpoints/model_management_endpoints.py index 1b49e2455e4..ec98d7d65f1 100644 --- a/litellm/proxy/management_endpoints/model_management_endpoints.py +++ b/litellm/proxy/management_endpoints/model_management_endpoints.py @@ -30,6 +30,7 @@ from litellm.litellm_core_utils.ptu_pricing import ( PTU_ZEROED_PRICING_FIELDS, PTU_ZEROED_TABLE_FIELDS, SEARCH_CONTEXT_SIZES, + ptu_config_error, ) from litellm.proxy._types import ( BlockModelRequest, @@ -308,42 +309,13 @@ def _raise_if_ptu_cost_attribution_disabled(incoming_model_info: Mapping[str, ob def _validate_ptu_model_info(model_info: Mapping[str, object]) -> None: """Enforce the PTU cross-field invariant on the effective model_info. - ptu_count and cost_per_ptu_per_hour must be set together, and a team_id and a - ptu_effective_from are required when they are. The start is mandatory rather than - defaulted because flat cost accrues from it: inferring one would let a deployment - configured today be billed for days it did not exist. Per-field bounds (positive - count, non-negative rate) are enforced by ModelInfo itself. - - Window ordering is checked before the count/rate gate. A patch that touches only one - end of the window carries no count or rate, and ModelInfo sees one field at a time, so - leaving it to either would let an inverted window reach the row; the next load then - fails to parse it and drops the deployment out of the router, where no further patch - can repair it because each one re-parses the stored value first. + The rules live in litellm_core_utils.ptu_pricing so that config.yaml registration + refuses the same deployments this endpoint does, for the same reason. Per-field bounds + (positive count, non-negative rate) are enforced by ModelInfo itself. """ - effective_from: Final = _coerce_ptu_datetime(model_info.get("ptu_effective_from")) - effective_to: Final = _coerce_ptu_datetime(model_info.get("ptu_effective_to")) - if effective_from is not None and effective_to is not None and effective_to <= effective_from: - raise HTTPException(status_code=400, detail="ptu_effective_to must be after ptu_effective_from") - - has_count: Final = model_info.get("ptu_count") is not None - has_rate: Final = model_info.get("cost_per_ptu_per_hour") is not None - if not has_count and not has_rate: - return - if has_count != has_rate: - raise HTTPException(status_code=400, detail="ptu_count and cost_per_ptu_per_hour must be set together") - if effective_from is None: - raise HTTPException( - status_code=400, - detail=( - "ptu_effective_from is required when PTU fields are set. Flat cost accrues from that " - "instant, so without it the start would have to be inferred and a deployment configured " - "today could be billed for days it did not exist" - ), - ) - if not model_info.get("team_id"): - raise HTTPException( - status_code=400, detail="team_id is required when PTU fields are set (one model maps to one team)" - ) + error: Final = ptu_config_error(model_info) + if error is not None: + raise HTTPException(status_code=400, detail=error) # The mirrored per-token pricing fields plus the three remaining fields @@ -515,28 +487,6 @@ def _ptu_priced_deployment(model_params: Deployment) -> Deployment: ) -def _parse_ptu_datetime(value: object) -> datetime.datetime | None: - """``value`` as a datetime, parsing an ISO string, else None.""" - if isinstance(value, datetime.datetime): - return value - if not isinstance(value, str): - return None - try: - return datetime.datetime.fromisoformat(value.replace("Z", "+00:00")) - except ValueError: - return None - - -def _coerce_ptu_datetime(value: object) -> datetime.datetime | None: - """Coerce a model_info effective-window value (datetime or ISO string) to UTC, else None.""" - parsed: Final = _parse_ptu_datetime(value) - if parsed is None: - return None - if parsed.tzinfo is None: - return parsed.replace(tzinfo=datetime.timezone.utc) - return parsed.astimezone(datetime.timezone.utc) - - def update_db_model(db_model: Deployment, updated_patch: updateDeployment) -> PrismaCompatibleUpdateDBModel: if updated_patch.model_info is not None: _raise_if_ptu_cost_attribution_disabled(updated_patch.model_info.model_dump(exclude_none=True)) diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index e50e9bf0537..8efa2c06998 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -20,7 +20,7 @@ from fastapi import ( Response, ) from pydantic import BaseModel, TypeAdapter, ValidationError -from typing_extensions import TypedDict, assert_never +from typing_extensions import ReadOnly, TypedDict, assert_never import litellm from litellm._logging import verbose_proxy_logger @@ -164,6 +164,12 @@ class UserProvisionerHelpers: """ Check if a user with the given email already exists and update them if found. + The matched row keeps its existing user_id even when the SCIM userName differs. + Virtual keys, team rosters, team/organization memberships and spend logs all + reference that id, so re-keying the user row would strand every one of them and + make removals against rosters holding the old id no-op. SCIM ids are opaque to + the client, which reads the stable id back from the response. + When admin_group is configured the resolved global role on new_user_request is persisted too, so re-upserting an existing email demotes a user who is no longer in the admin group instead of leaving the stale role. @@ -189,20 +195,22 @@ class UserProvisionerHelpers: new_teams: Final = list(dict.fromkeys(new_user_request.teams or [])) if new_user_request.user_id != existing_user.user_id: - await _table(UserRepository(prisma_client)).update( - where={"user_id": existing_user.user_id}, - data={"user_id": new_user_request.user_id}, + verbose_proxy_logger.info( + "SCIM: email %s already provisioned as user_id=%s, keeping that id instead of re-keying to %s", + new_user_request.user_email, + existing_user.user_id, + new_user_request.user_id, ) await _handle_team_membership_changes( - user_id=new_user_request.user_id, + user_id=existing_user.user_id, existing_teams=existing_user.teams or [], new_teams=new_teams, raise_on_error=True, ) updated_user: Final = await _table(UserRepository(prisma_client)).update( - where={"user_id": new_user_request.user_id}, + where={"user_id": existing_user.user_id}, data={ "user_email": new_user_request.user_email, "user_alias": new_user_request.user_alias, @@ -619,6 +627,40 @@ def _admitted_member_ids(classified: Iterable[_ClassifiedGroupMember], created_i ) +class _UserIdWhere(TypedDict): + user_id: ReadOnly[str] + + +class _ScimErrorDetail(TypedDict): + error: ReadOnly[str] + + +async def _ensure_group_member_user( + user_id: str, + created_via: str, + prisma_client: PrismaClient, +) -> NewUserResponse | None: + """The created user, or None when the id already resolves to a user row (a + concurrent provisioning request won the creation race after our lookup missed). + + Raises: + HTTPException: 500 when the user can neither be created nor found. The + request has to fail so the identity provider retries, instead of recording + success for a member the roster silently dropped. + """ + created: Final = await _create_user_if_not_exists(user_id=user_id, created_via=created_via) + if created is not None: + return created + where: Final[_UserIdWhere] = {"user_id": user_id} + existing: Final = await _table(UserRepository(prisma_client)).find_unique(where=where) + if existing is not None: + return None + detail: Final[_ScimErrorDetail] = { + "error": f"Failed to create user '{user_id}' while provisioning group membership." + } + raise HTTPException(status_code=500, detail=detail) + + async def _resolve_group_member_ids( members: Sequence[SCIMMember], created_via: str, @@ -636,7 +678,8 @@ async def _resolve_group_member_ids( Raises: HTTPException: 400 when a member id is empty, or when scim_upsert_user is False and a member id is neither an existing user, an existing team, nor a - member declared to be something other than a user. + member declared to be something other than a user. 500 when a member's + user row can neither be created nor found. """ classified: Final = tuple([await _classify_group_member(member, prisma_client) for member in members]) partition: Final = _partition_classified_members(classified) @@ -657,10 +700,14 @@ async def _resolve_group_member_ids( }, ) + unique_unknown_ids: Final = tuple(dict.fromkeys(partition.unknown_ids)) creations: Final = tuple( [ - (user_id, await _create_user_if_not_exists(user_id=user_id, created_via=created_via)) - for user_id in partition.unknown_ids + ( + user_id, + await _ensure_group_member_user(user_id=user_id, created_via=created_via, prisma_client=prisma_client), + ) + for user_id in unique_unknown_ids ] ) created_users: Final = tuple(created for _, created in creations if created is not None) @@ -668,10 +715,7 @@ async def _resolve_group_member_ids( return GroupMemberExtractionResult( existing_member_ids=partition.resolved_ids, created_users=created_users, - all_member_ids=_admitted_member_ids( - classified, - frozenset(user_id for user_id, created in creations if created is not None), - ), + all_member_ids=_admitted_member_ids(classified, frozenset(unique_unknown_ids)), ) @@ -2371,19 +2415,25 @@ async def _apply_group_patch_updates(group_id: str, update_data: dict[str, objec async def _handle_group_membership_changes(group_id: str, current_members: set[str], final_members: set[str]): - """Handle adding/removing members from the group.""" + """Handle adding/removing members from the group. + + Runs strict: a genuine add or remove failure propagates so the group request + fails and the identity provider retries, instead of reporting success for a + member the roster never received. Idempotent no-ops (already in / already out + of the team) are still swallowed by patch_team_membership. + """ members_to_add: Final = final_members - current_members members_to_remove: Final = current_members - final_members verbose_proxy_logger.debug("members_to_add: %s", members_to_add) verbose_proxy_logger.debug("members_to_remove: %s", members_to_remove) - # Use existing helper functions for team membership changes for member_id in members_to_add: await patch_team_membership( user_id=member_id, teams_ids_to_add_user_to=[group_id], teams_ids_to_remove_user_from=[], + raise_on_error=True, ) for member_id in members_to_remove: @@ -2391,6 +2441,7 @@ async def _handle_group_membership_changes(group_id: str, current_members: set[s user_id=member_id, teams_ids_to_add_user_to=[], teams_ids_to_remove_user_from=[group_id], + raise_on_error=True, ) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 1ebcb53fd6b..3c135650de9 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -19,6 +19,7 @@ import secrets from collections.abc import Mapping, Sequence from copy import deepcopy from html import escape +from types import MappingProxyType from typing import ( TYPE_CHECKING, Annotated, @@ -245,6 +246,7 @@ def _team_detail_db(repo: "_HasTeamDetailTable") -> "_PrismaTableActions[_TeamDe _MODEL_ALIASES_ADAPTER: Final = TypeAdapter(dict[str, str]) +_SSO_TOKEN_CLAIMS_ADAPTER: Final = TypeAdapter(Mapping[str, object]) def _decode_model_aliases(value: object) -> object: @@ -1002,6 +1004,30 @@ def process_sso_jwt_access_token( return None +def _decode_sso_token_claims(token: str | None) -> Mapping[str, object]: + if not token: + return MappingProxyType({}) + try: + return MappingProxyType( + _SSO_TOKEN_CLAIMS_ADAPTER.validate_python(jwt.decode(token, options={"verify_signature": False})) + ) + except (jwt.exceptions.InvalidTokenError, ValidationError): + verbose_proxy_logger.debug("SSO token is not a decodable JWT, skipping token claims") + return MappingProxyType({}) + + +def _merge_sso_token_claims( + userinfo: Mapping[str, object], + id_token: str | None, + access_token: str | None, +) -> Mapping[str, object]: + sources: Final = (userinfo, _decode_sso_token_claims(id_token), _decode_sso_token_claims(access_token)) + claim_names: Final = frozenset(key for source in sources for key in source) + return MappingProxyType( + {key: next((source[key] for source in sources if source.get(key) is not None), None) for key in claim_names} + ) + + async def _raise_if_sso_exceeds_free_user_limit(premium_user: bool, prisma_client: PrismaClient | None) -> None: """Free tier allows SSO for up to 5 billable users; beyond that requires an Enterprise license.""" if premium_user is True: @@ -1534,12 +1560,34 @@ async def get_generic_sso_response( role_mappings: Final = await _setup_role_mappings() team_mappings: Final = await _setup_team_mappings() + generic_include_token_claims: Final = os.getenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "false").lower() == "true" - def response_convertor(response, client): + def response_convertor(response: Mapping[str, object], httpx_session: object): nonlocal received_response # return for user debugging - received_response = response + response_id_token: Final = response.get("id_token") + response_access_token: Final = response.get("access_token") + id_token: Final = ( + response_id_token if isinstance(response_id_token, str) and response_id_token else generic_sso.id_token + ) + access_token: Final = ( + response_access_token + if isinstance(response_access_token, str) and response_access_token + else generic_sso.access_token + ) + claims: Final = ( + _merge_sso_token_claims( + userinfo=response, + id_token=id_token, + access_token=access_token, + ) + if generic_include_token_claims + else response + ) + received_response = { # mutable-ok: preserve the existing dict return contract + key: value for key, value in claims.items() if key not in _OAUTH_TOKEN_FIELDS + } return generic_response_convertor( - response=response, + response=claims, jwt_handler=jwt_handler, sso_jwt_handler=sso_jwt_handler, role_mappings=role_mappings, @@ -1641,13 +1689,6 @@ async def get_generic_sso_response( # Pass the full response so custom response_convertor implementations # can access all fields (including id_token for claim extraction). result = response_convertor(combined_response, generic_sso) - # Strip bearer credentials from combined_response before storing in - # received_response. received_response may appear in restricted-group - # error messages — bearer tokens (access_token, id_token, refresh_token) - # must not be exposed to callers. - # Assign directly rather than relying on nonlocal mutation so that Pyright - # can track that received_response is non-None from this point on. - received_response = {k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS} sso_assertion = assertion_from_sso_login( combined_response.get("id_token"), combined_response.get("refresh_token") ) diff --git a/litellm/proxy/prisma_migration.py b/litellm/proxy/prisma_migration.py index 6f9561afec9..373c3811949 100644 --- a/litellm/proxy/prisma_migration.py +++ b/litellm/proxy/prisma_migration.py @@ -1,26 +1,44 @@ -# What is this? -## Script to apply initial prisma migration on Docker setup +"""Standalone entrypoint for applying database migrations and generating the Prisma client. + +The entrypoint enforces migration failures by default. Set +ENFORCE_PRISMA_MIGRATION_CHECK=false to preserve log-only behavior for migration and +Prisma generate failures. +""" import os import subprocess import sys -sys.path.insert(0, os.path.abspath("./")) # Adds the parent directory to the system path +sys.path.insert(0, os.path.abspath("./")) from typing import Final from litellm._logging import verbose_proxy_logger from litellm.proxy.proxy_cli import run_server +from litellm.secret_managers.main import str_to_bool -# Call the Click command with standalone_mode=False -run_server(["--skip_server_startup"], standalone_mode=False) -# run prisma generate -verbose_proxy_logger.info("Running 'prisma generate'...") -result: Final = subprocess.run(["prisma", "generate"], capture_output=True, text=True) -verbose_proxy_logger.info("'prisma generate' stdout: %s", result.stdout) # Log stdout -exit_code: Final = result.returncode +def main() -> int: + enforce_prisma_migration_check: Final = str_to_bool(os.getenv("ENFORCE_PRISMA_MIGRATION_CHECK")) is not False + run_server_args: Final = ( + ("--skip_server_startup", "--enforce_prisma_migration_check") + if enforce_prisma_migration_check + else ("--skip_server_startup",) + ) + run_server(run_server_args, standalone_mode=False) -if exit_code != 0: - verbose_proxy_logger.info("'prisma generate' failed with exit code %s.", exit_code) - verbose_proxy_logger.error("'prisma generate' stderr: %s", result.stderr) # Log stderr + verbose_proxy_logger.info("Running 'prisma generate'...") + result: Final = subprocess.run(("prisma", "generate"), capture_output=True, text=True) + verbose_proxy_logger.info("'prisma generate' stdout: %s", result.stdout) + exit_code: Final = result.returncode + + if exit_code != 0: + verbose_proxy_logger.info("'prisma generate' failed with exit code %s.", exit_code) + verbose_proxy_logger.error("'prisma generate' stderr: %s", result.stderr) + if enforce_prisma_migration_check: + return exit_code + return 0 + + +if __name__ == "__main__": + sys.exit(main()) diff --git a/litellm/proxy/proxy_cli.py b/litellm/proxy/proxy_cli.py index 6a0b3c6bfb2..0e3e43accef 100644 --- a/litellm/proxy/proxy_cli.py +++ b/litellm/proxy/proxy_cli.py @@ -19,7 +19,6 @@ from pydantic import BaseModel, ConfigDict import litellm from litellm.constants import DEFAULT_NUM_WORKERS_LITELLM_PROXY from litellm.proxy.db.query_engine_reaper import start_query_engine_reaper -from litellm.secret_managers.main import get_secret_bool if TYPE_CHECKING: from fastapi import FastAPI @@ -790,6 +789,12 @@ class ProxyInitializationHelpers: is_flag=True, help="Connects to RDS DB with IAM token", ) +@click.option( + "--azure_postgresql_auth", + default=False, + is_flag=True, + help="Connects to Azure Database for PostgreSQL with a Microsoft Entra ID token", +) @click.option( "--num_requests", default=10, @@ -951,6 +956,7 @@ def run_server( granian_threads, test_async, iam_token_db_auth, + azure_postgresql_auth: bool, num_requests, use_queue, health, @@ -1080,31 +1086,27 @@ def run_server( db_statement_timeout: float | None = None db_lock_timeout: float | None = None general_settings = {} - ### GET DB TOKEN FOR IAM AUTH ### + ### GET DB TOKEN FOR RDS IAM / AZURE ENTRA AUTH ### - if iam_token_db_auth or get_secret_bool("IAM_TOKEN_DB_AUTH"): - from litellm.proxy.auth.rds_iam_token import generate_iam_auth_token + from litellm.proxy.db.db_url_settings import DatabaseURLSettings + from litellm.proxy.db.token_auth import ( + AZURE_POSTGRESQL_AUTH_ENV_VAR, + IAM_TOKEN_DB_AUTH_ENV_VAR, + token_auth_flag_enabled, + ) - db_host: Final = os.getenv("DATABASE_HOST") - # Default to the Postgres standard port. Without a default, - # `db_port=None` flows into `boto.generate_db_auth_token(Port=None)` - # and botocore stringifies it to `"None"` while building the - # presigned URL, which then blows up with `ValueError: Port could - # not be cast to integer value as 'None'` during signing. - db_port: Final = os.getenv("DATABASE_PORT", "5432") - db_user: Final = os.getenv("DATABASE_USER") - db_name: Final = os.getenv("DATABASE_NAME") - db_schema: Final = os.getenv("DATABASE_SCHEMA") - - token: Final = generate_iam_auth_token(db_host=db_host, db_port=db_port, db_user=db_user) - - # print(f"token: {token}") - _db_url = f"postgresql://{db_user}:{token}@{db_host}:{db_port}/{db_name}" - if db_schema: - _db_url += f"?schema={db_schema}" - - os.environ["DATABASE_URL"] = _db_url - os.environ["IAM_TOKEN_DB_AUTH"] = "True" + wants_rds_iam: Final = iam_token_db_auth or token_auth_flag_enabled( + os.getenv(IAM_TOKEN_DB_AUTH_ENV_VAR), env_var=IAM_TOKEN_DB_AUTH_ENV_VAR + ) + wants_azure_entra: Final = azure_postgresql_auth or token_auth_flag_enabled( + os.getenv(AZURE_POSTGRESQL_AUTH_ENV_VAR), env_var=AZURE_POSTGRESQL_AUTH_ENV_VAR + ) + if wants_rds_iam: + os.environ[IAM_TOKEN_DB_AUTH_ENV_VAR] = "True" + if wants_azure_entra: + os.environ[AZURE_POSTGRESQL_AUTH_ENV_VAR] = "True" + if wants_rds_iam or wants_azure_entra: + DatabaseURLSettings.from_env().apply_writer_url_to_env() ### DECRYPT ENV VAR ### diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 12d0f13ebdd..4b97cade7f0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -253,6 +253,7 @@ from litellm.exceptions import RejectedRequestError from litellm.integrations.custom_guardrail import ModifyResponseException from litellm.integrations.custom_logger import CustomLogger from litellm.integrations.SlackAlerting.slack_alerting import SlackAlerting +from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, get_litellm_metadata_from_kwargs, @@ -2994,6 +2995,13 @@ async def _is_spend_counter_cache_warm(counter_key: str) -> bool: return spend_counter_cache.in_memory_cache.get_cache(key=counter_key) is not None +async def increment_spend_counter(counter_key: str, increment: float): + """Public raw-counter increment for budget domains outside the entity scopes (e.g. + shadow eval's per-leg spend), sharing the primitive the entity counters use so + invalidation and read semantics can never drift.""" + return await _increment_spend_counter_cache(counter_key=counter_key, increment=increment) + + async def _increment_spend_counter_cache(counter_key: str, increment: float): if spend_counter_cache.redis_cache is not None: try: @@ -6308,7 +6316,8 @@ class ProxyConfig: # Schedule new job if retention period is set (not None) retention_period: Final = general_settings.get("maximum_spend_logs_retention_period") autorouter_retention: Final = general_settings.get("maximum_autorouter_session_retention_period") - if retention_period is not None or autorouter_retention is not None: + health_check_retention: Final = general_settings.get("maximum_health_check_retention_period") + if retention_period is not None or autorouter_retention is not None or health_check_retention is not None: from litellm.proxy.db.db_transaction_queue.spend_log_cleanup import ( SpendLogCleanup, ) @@ -6463,6 +6472,13 @@ class ProxyConfig: if old_session_value != new_session_value: await self._reschedule_spend_log_cleanup_job() + if "maximum_health_check_retention_period" in _general_settings: + old_health_check_value: Final = general_settings.get("maximum_health_check_retention_period") + new_health_check_value: Final = _general_settings["maximum_health_check_retention_period"] + general_settings["maximum_health_check_retention_period"] = new_health_check_value + if old_health_check_value != new_health_check_value: + await self._reschedule_spend_log_cleanup_job() + ## SPEND LOG CLEANUP BOUNDS ## # The dashboard writes these straight to the DB, so without copying them # here the running cleanup job never sees them. A key the DB no longer @@ -9078,6 +9094,7 @@ class ProxyStartupEvent: if ( general_settings.get("maximum_spend_logs_retention_period") is not None or general_settings.get("maximum_autorouter_session_retention_period") is not None + or general_settings.get("maximum_health_check_retention_period") is not None ): spend_log_cleanup: Final = SpendLogCleanup() cleanup_cron: Final = general_settings.get("maximum_spend_logs_cleanup_cron") @@ -11909,8 +11926,6 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) Returns: TokenCountResponse """ - from litellm import token_counter - global llm_router prompt: Final = request.prompt @@ -11994,7 +12009,7 @@ async def token_counter(request: TokenCountRequest, call_endpoint: bool = False) _tokenizer_used: Final = litellm.utils._select_tokenizer(model=model_to_use, custom_tokenizer=custom_tokenizer) tokenizer_used: Final = str(_tokenizer_used["type"]) - total_tokens: Final = token_counter( + total_tokens: Final = await asyncify(litellm.token_counter)( model=model_to_use, text=prompt, messages=messages, @@ -15813,6 +15828,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro "store_model_in_db": "Boolean", "store_prompts_in_spend_logs": "Boolean", "maximum_spend_logs_retention_period": "String", + "maximum_health_check_retention_period": "String", "maximum_spend_logs_cleanup_batch_size": "Integer", "maximum_spend_logs_cleanup_max_batches": "Integer", "maximum_spend_logs_cleanup_run_budget": "String", diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 60058c777ca..d9959677116 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1502,7 +1502,8 @@ model LiteLLM_ShadowEvalJob { baseline_model String? // reverse only: the fixed model the router is judged against judge_model String shadow_percentage Float - max_turns Int // this key's sample budget: judge at most this many turns + max_turns Int // sample-count ceiling: the whole budget on pre-max_budget jobs, the error-loop valve otherwise + max_budget Float? // per-key USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets created_at DateTime @default(now()) created_by String? ends_at DateTime @@ -1525,6 +1526,7 @@ model LiteLLM_ShadowEvalAttempt { shadow_model String? confidence Float? judge_cost Float @default(0) + shadow_cost Float @default(0) error String? created_at DateTime @default(now()) diff --git a/litellm/proxy/spend_tracking/budget_reservation.py b/litellm/proxy/spend_tracking/budget_reservation.py index 17074ec967b..ce6c9330620 100644 --- a/litellm/proxy/spend_tracking/budget_reservation.py +++ b/litellm/proxy/spend_tracking/budget_reservation.py @@ -5,6 +5,7 @@ import json from collections.abc import Mapping, Sequence from dataclasses import dataclass from datetime import datetime, timedelta, timezone +from types import MappingProxyType from typing import Any, Final, NoReturn, cast from fastapi import HTTPException, status @@ -182,11 +183,18 @@ async def reserve_budget_for_request( if not counters: return None + input_token_counts: Final = await count_request_input_tokens( + request_body=request_body, + route=route, + llm_router=llm_router, + ) + current_spend_by_counter_key: Final[dict[str, float]] = {} reservation_cost = estimate_request_max_cost( request_body=request_body, route=route, llm_router=llm_router, + input_token_counts=input_token_counts, ) # estimate_request_max_cost still returns None when the model is unknown # to the cost map (no token-priced cost fields, e.g. image/audio routes). @@ -245,7 +253,12 @@ async def reserve_budget_for_request( if not applied_entries: return None - input_cost: Final = estimate_request_input_cost(request_body=request_body, route=route, llm_router=llm_router) + input_cost: Final = estimate_request_input_cost( + request_body=request_body, + route=route, + llm_router=llm_router, + input_token_counts=input_token_counts, + ) return { "reserved_cost": reservation_cost, "entries": applied_entries, @@ -907,20 +920,17 @@ def estimate_request_max_cost( request_body: dict, route: str, llm_router: Router | None, + input_token_counts: Mapping[str, int] | None = None, ) -> float | None: - model: Final = get_model_from_request(request_body, route, llm_router=llm_router) - if model is None: - return None - - models: Final = [model] if isinstance(model, str) else model estimates = [ _estimate_request_max_cost_for_model( request_body=request_body, route=route, model=model_name, llm_router=llm_router, + input_tokens=(input_token_counts or {}).get(model_name), ) - for model_name in models + for model_name in _get_request_models(request_body=request_body, route=route, llm_router=llm_router) ] estimates = [estimate for estimate in estimates if estimate is not None] if not estimates: @@ -932,6 +942,7 @@ def estimate_request_input_cost( request_body: dict, route: str, llm_router: Router | None, + input_token_counts: Mapping[str, int] | None = None, ) -> float | None: """Cost of the request's input tokens alone. @@ -940,19 +951,15 @@ def estimate_request_input_cost( cancelled in-flight request has already incurred. A cancelled reservation is reconciled to this instead of being refunded to zero. """ - model: Final = get_model_from_request(request_body, route, llm_router=llm_router) - if model is None: - return None - - models: Final = [model] if isinstance(model, str) else model estimates = [ _estimate_request_input_cost_for_model( request_body=request_body, route=route, model=model_name, llm_router=llm_router, + input_tokens=(input_token_counts or {}).get(model_name), ) - for model_name in models + for model_name in _get_request_models(request_body=request_body, route=route, llm_router=llm_router) ] estimates = [estimate for estimate in estimates if estimate is not None] if not estimates: @@ -965,6 +972,7 @@ def _estimate_request_input_cost_for_model( route: str, model: str, llm_router: Router | None, + input_tokens: int | None = None, ) -> float | None: estimates: Final = [ _input_cost_for_cost_info( @@ -972,6 +980,7 @@ def _estimate_request_input_cost_for_model( route=route, model=model, model_info=model_info, + input_tokens=input_tokens, ) for model_info in _get_model_cost_infos(model=model, llm_router=llm_router) ] @@ -984,24 +993,26 @@ def _input_cost_for_cost_info( route: str, model: str, model_info: Mapping[str, object], + input_tokens: int | None = None, ) -> float | None: - input_tokens: Final = _estimate_input_tokens( + estimated_input_tokens: Final = _estimate_input_tokens( request_body=request_body, route=route, model=model, model_info=model_info, + input_tokens=input_tokens, ) - if input_tokens is None: + if estimated_input_tokens is None: return None tiered_pricing: Final = model_info.get("tiered_pricing") if isinstance(tiered_pricing, list) and tiered_pricing: - tier: Final = select_tier_for_input(tiered_pricing=tiered_pricing, input_tokens=input_tokens) + tier: Final = select_tier_for_input(tiered_pricing=tiered_pricing, input_tokens=estimated_input_tokens) if tier is not None: - return input_tokens * tier_rate(tier, "input_cost_per_token") + return estimated_input_tokens * tier_rate(tier, "input_cost_per_token") input_cost_per_token: Final = _to_float(model_info.get("input_cost_per_token")) if input_cost_per_token is None: return None - return input_tokens * input_cost_per_token + return estimated_input_tokens * input_cost_per_token def _estimate_request_max_cost_for_model( @@ -1009,6 +1020,7 @@ def _estimate_request_max_cost_for_model( route: str, model: str, llm_router: Router | None, + input_tokens: int | None = None, ) -> float | None: estimates: Final = [ _max_cost_for_cost_info( @@ -1016,6 +1028,7 @@ def _estimate_request_max_cost_for_model( route=route, model=model, model_info=model_info, + input_tokens=input_tokens, ) for model_info in _get_model_cost_infos(model=model, llm_router=llm_router) ] @@ -1028,6 +1041,7 @@ def _max_cost_for_cost_info( route: str, model: str, model_info: Mapping[str, object], + input_tokens: int | None = None, ) -> float | None: image_cost: Final = _estimate_image_generation_cost( request_body=request_body, @@ -1036,30 +1050,31 @@ def _max_cost_for_cost_info( if image_cost is not None: return image_cost - input_tokens: Final = _estimate_input_tokens( + estimated_input_tokens: Final = _estimate_input_tokens( request_body=request_body, route=route, model=model, model_info=model_info, + input_tokens=input_tokens, ) output_tokens: Final = _estimate_output_tokens( request_body=request_body, route=route, model_info=model_info, ) - if input_tokens is None or output_tokens is None: + if estimated_input_tokens is None or output_tokens is None: return None output_multiplier: Final = _get_output_multiplier(request_body=request_body) tiered_pricing: Final = model_info.get("tiered_pricing") if isinstance(tiered_pricing, list) and tiered_pricing: - tier: Final = select_tier_for_input(tiered_pricing=tiered_pricing, input_tokens=input_tokens) + tier: Final = select_tier_for_input(tiered_pricing=tiered_pricing, input_tokens=estimated_input_tokens) if tier is not None: output_rate = max( tier_rate(tier, "output_cost_per_token"), tier_rate(tier, "output_cost_per_reasoning_token"), ) - return (input_tokens * tier_rate(tier, "input_cost_per_token")) + ( + return (estimated_input_tokens * tier_rate(tier, "input_cost_per_token")) + ( output_tokens * output_multiplier * output_rate ) @@ -1068,8 +1083,8 @@ def _max_cost_for_cost_info( output_cost_per_reasoning_token: Final = _to_float(model_info.get("output_cost_per_reasoning_token")) cost = 0.0 if input_cost_per_token is not None: - cost += input_tokens * input_cost_per_token - elif input_tokens > 0: + cost += estimated_input_tokens * input_cost_per_token + elif estimated_input_tokens > 0: return None # The reasoning-token share is unknown before the request runs, so reserve every @@ -1192,12 +1207,70 @@ def _get_deployment_tiered_pricing_tables( ] -def _estimate_input_tokens( +def _get_request_models( request_body: dict, route: str, - model: str, - model_info: Mapping[str, object], -) -> int | None: + llm_router: Router | None, +) -> Sequence[str]: + model: Final = get_model_from_request(request_body, route, llm_router=llm_router) + if model is None: + return () + return (model,) if isinstance(model, str) else tuple(model) + + +TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS: Final = 30_000 + + +async def count_request_input_tokens( + request_body: dict, + route: str, + llm_router: Router | None, +) -> Mapping[str, int]: + """Input-token count per candidate model, counted once per request. + + Tokenizing is the reservation path's dominant CPU cost and is O(prompt), so + counting a large prompt inline stalls every other request on the worker. + Large prompts are counted in a worker thread, and the counts are reused by + both the max-cost and the input-cost estimate. + """ + models: Final = _get_request_models(request_body=request_body, route=route, llm_router=llm_router) + if not models: + return MappingProxyType({}) + if _approximate_input_size(request_body) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS: + return _count_input_tokens_for_models(request_body=request_body, models=models) + return await asyncio.to_thread( + _count_input_tokens_for_models, + request_body=request_body, + models=models, + ) + + +def _count_input_tokens_for_models( + request_body: dict, + models: Sequence[str], +) -> Mapping[str, int]: + return MappingProxyType( + { + model: tokens + for model in models + if (tokens := _count_input_tokens(request_body=request_body, model=model)) is not None + } + ) + + +_INPUT_SIZE_FIELDS: Final = ("messages", "prompt", "input", "query", "documents", "tools", "tool_choice") + + +def _approximate_input_size(request_body: dict) -> int: + """Length of the request's input text, a cheap stand-in for tokenizing cost. + + Every field _count_input_tokens hands the tokenizer is sized here, and + rendering rather than walking keeps mapping keys in the total, which a tool + schema's property names are.""" + return sum(len(str(request_body.get(field, ""))) for field in _INPUT_SIZE_FIELDS) + + +def _count_input_tokens(request_body: dict, model: str) -> int | None: try: if "messages" in request_body: return litellm.token_counter( @@ -1219,6 +1292,21 @@ def _estimate_input_tokens( return query_tokens + document_tokens except Exception: verbose_proxy_logger.debug("Unable to count input tokens for budget reservation", exc_info=True) + return None + + +def _estimate_input_tokens( + request_body: dict, + route: str, + model: str, + model_info: Mapping[str, object], + input_tokens: int | None = None, +) -> int | None: + counted: Final = ( + input_tokens if input_tokens is not None else _count_input_tokens(request_body=request_body, model=model) + ) + if counted is not None: + return counted max_input_tokens: Final = _to_int(model_info.get("max_input_tokens")) if max_input_tokens is not None: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 58e6f5a94e2..81a86ebe34d 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -127,6 +127,11 @@ from litellm.proxy.db.spend_log_batching import ( spend_log_row_bytes, spend_log_write_batches, ) +from litellm.proxy.db.token_auth import ( + DatabaseTokenAuth, + mint_database_token, + resolve_database_token_auth, +) from litellm.proxy.guardrails.guardrail_hooks.unified_guardrail.unified_guardrail import ( UnifiedLLMGuardrails, ) @@ -3304,7 +3309,7 @@ class PrismaClient: ): ## init logging object self.proxy_logging_obj = proxy_logging_obj - self.iam_token_db_auth: bool | None = str_to_bool(os.getenv("IAM_TOKEN_DB_AUTH")) + self.token_auth: DatabaseTokenAuth | None = resolve_database_token_auth() verbose_proxy_logger.debug("Creating Prisma Client..") try: from prisma import Prisma @@ -3313,22 +3318,22 @@ class PrismaClient: verbose_proxy_logger.error("This usually means 'prisma generate' hasn't been run yet.") verbose_proxy_logger.error("Please run 'prisma generate' to generate the Prisma client.") raise Exception("Unable to find Prisma binaries. Please run 'prisma generate' first.") - iam_flag: Final = self.iam_token_db_auth if self.iam_token_db_auth is not None else False + token_auth: Final = self.token_auth # When read-replica routing is on, tag log lines with [writer]/[reader] - # so the two wrappers' interleaved IAM refresh logs can be told apart. + # so the two wrappers' interleaved token refresh logs can be told apart. # Single-DB deployments get an empty prefix (logs unchanged). read_replica_url = os.getenv("DATABASE_URL_READ_REPLICA") writer_log_prefix: Final = "[writer]" if read_replica_url else "" if http_client is not None: writer_wrapper = PrismaWrapper( original_prisma=Prisma(http=http_client), - iam_token_db_auth=iam_flag, + token_auth=token_auth, log_prefix=writer_log_prefix, ) else: writer_wrapper = PrismaWrapper( original_prisma=Prisma(), - iam_token_db_auth=iam_flag, + token_auth=token_auth, log_prefix=writer_log_prefix, ) @@ -3340,29 +3345,22 @@ class PrismaClient: self.db: PrismaWrapper | RoutingPrismaWrapper if read_replica_url: try: - # If IAM auth is enabled, the reader refreshes its own token on + # If token auth is enabled, the reader refreshes its own token on # the same cadence as the writer. We parse the static endpoint # pieces (host/port/user/db) once from the reader URL — only - # the IAM token rotates after that. - reader_iam_endpoint: Final = parse_iam_endpoint_from_url(read_replica_url) if iam_flag else None - # Mint a fresh IAM token for the reader BEFORE constructing the + # the token rotates after that. + reader_iam_endpoint: Final = ( + parse_iam_endpoint_from_url(read_replica_url) if token_auth is not None else None + ) + # Mint a fresh token for the reader BEFORE constructing the # Prisma client. Mirrors what `proxy_cli.py` already does for - # the writer (proxy_cli.py:812-832) — without this, the reader - # Prisma is built with whatever placeholder URL the user - # supplied (no real token), and the first query falls through - # to the synchronous fallback path in - # `PrismaWrapper.__getattr__`, which deadlocks the event loop - # and times out after 30s. - if iam_flag and reader_iam_endpoint is not None: - from litellm.proxy.auth.rds_iam_token import ( - generate_iam_auth_token, - ) - - reader_token: Final = generate_iam_auth_token( - db_host=reader_iam_endpoint.host, - db_port=reader_iam_endpoint.port, - db_user=reader_iam_endpoint.user, - ) + # the writer — without this, the reader Prisma is built with + # whatever placeholder URL the user supplied (no real token), + # and the first query falls through to the synchronous fallback + # path in `PrismaWrapper.__getattr__`, which deadlocks the event + # loop and times out after 30s. + if token_auth is not None and reader_iam_endpoint is not None: + reader_token: Final = mint_database_token(token_auth, reader_iam_endpoint) read_replica_url = reader_iam_endpoint.build_url(reader_token) os.environ["DATABASE_URL_READ_REPLICA"] = read_replica_url reader_kwargs: Final[dict[str, Any]] = {"datasource": {"url": read_replica_url}} @@ -3372,7 +3370,7 @@ class PrismaClient: reader_prisma = Prisma(**reader_kwargs) reader_wrapper: Final = PrismaWrapper( original_prisma=reader_prisma, - iam_token_db_auth=iam_flag, + token_auth=token_auth, db_url_env_var="DATABASE_URL_READ_REPLICA", iam_endpoint=reader_iam_endpoint, recreate_uses_datasource=True, @@ -3381,15 +3379,15 @@ class PrismaClient: self.db = RoutingPrismaWrapper(writer=writer_wrapper, reader=reader_wrapper) verbose_proxy_logger.info( "PrismaClient: read-replica routing enabled via DATABASE_URL_READ_REPLICA" - + (" (with IAM token auto-refresh)" if iam_flag else "") + + (f" (with {token_auth.label} auto-refresh)" if token_auth is not None else "") ) except Exception as e: # Reader is opt-in; never let its construction fail proxy # startup. Mirrors the runtime contract from # `RoutingPrismaWrapper.connect`: reader-side failures are # logged and we keep serving traffic via the writer alone. - # This recovers from transient AWS STS hiccups during the - # reader IAM token mint, malformed DATABASE_URL_READ_REPLICA, + # This recovers from transient credential-provider hiccups + # during the reader token mint, malformed DATABASE_URL_READ_REPLICA, # and Prisma construction errors. Operator restart is required # to retry read-routing once the underlying issue is resolved. verbose_proxy_logger.warning( diff --git a/litellm/router.py b/litellm/router.py index e4e3a857411..e9eeab53934 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -52,7 +52,7 @@ from litellm.constants import ( SESSION_DEPLOYMENT_AFFINITY_TTL_METADATA_KEY, ) from litellm.integrations.custom_logger import CustomLogger -from litellm.litellm_core_utils.asyncify import run_async_function +from litellm.litellm_core_utils.asyncify import asyncify, run_async_function from litellm.litellm_core_utils.core_helpers import ( _get_parent_otel_span_from_kwargs, coerce_token_limit, @@ -64,7 +64,11 @@ from litellm.litellm_core_utils.coroutine_checker import coroutine_checker from litellm.litellm_core_utils.credential_accessor import CredentialAccessor from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging -from litellm.litellm_core_utils.ptu_pricing import zeroed_ptu_pricing +from litellm.litellm_core_utils.ptu_pricing import ( + is_ptu_cost_attribution_enabled, + ptu_config_error, + zeroed_ptu_pricing, +) from litellm.litellm_core_utils.request_timeout_resolver import ( get_configured_request_timeout, ) @@ -7701,9 +7705,11 @@ class Router: - None: If the deployment is not active for the current environment (if 'supported_environments' is set in litellm_params) """ try: - zeroed_pricing: Final = ( - zeroed_ptu_pricing(_model_info, _litellm_params) if _model_info.get("db_model") is not True else None - ) + config_sourced: Final = _model_info.get("db_model") is not True + ptu_error: Final = ptu_config_error(_model_info, model_name=_model_name) if config_sourced else None + if ptu_error is not None and is_ptu_cost_attribution_enabled(): + raise ValueError(ptu_error) + zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None litellm_params: Final[LiteLLM_Params] = LiteLLM_Params( **( _litellm_params @@ -9192,7 +9198,7 @@ class Router: # get_model_info() hands back an lru_cache'd dict, so merge into a copy; unset # values are skipped or Deployment's None pricing defaults would erase the map's - merged_model_info: Final = copy.copy(model_info) + merged_model_info: Final = copy.deepcopy(model_info) if user_model_info: for key, value in user_model_info.items(): if value is not None: @@ -9243,7 +9249,7 @@ class Router: litellm_model_name_model_info: ModelInfo | None = None try: - custom_model_info = litellm.model_cost.get(model_id) + custom_model_info = copy.deepcopy(litellm.model_cost.get(model_id)) except Exception: pass @@ -9258,9 +9264,8 @@ class Router: base_model: Final = custom_model_info.get("base_model", None) if base_model is not None: ## update litellm model info with base model info - base_model_info: Final = litellm.get_model_info(model=base_model) + base_model_info: Final = copy.deepcopy(litellm.get_model_info(model=base_model)) if base_model_info is not None: - custom_model_info = custom_model_info or {} # Base model provides defaults, custom model info overrides custom_model_info = _update_dictionary( cast(dict, base_model_info), @@ -9276,13 +9281,13 @@ class Router: model_info = cast( ModelInfo, _update_dictionary( - cast(dict, litellm_model_name_model_info).copy(), + copy.deepcopy(cast(dict, litellm_model_name_model_info)), custom_model_info, ), ) elif litellm_model_name_model_info is not None: # (2) Built-in only — no custom pricing to merge - model_info = litellm_model_name_model_info + model_info = copy.deepcopy(litellm_model_name_model_info) elif custom_model_info is not None: # (3) Custom only — model not in built-in cost map yet # custom_model_info already includes base_model defaults at this point, if applicable @@ -10570,6 +10575,64 @@ class Router: return litellm.token_counter(messages=cast(list, input_messages)) # cast-ok: transformed chat messages raise ValueError("Either messages or input must be provided to count tokens") + def _deployment_max_input_tokens(self, model: str, deployment: Mapping[str, object]) -> int | None: + """The deployment's declared context window, or None when it declares none or cannot be resolved.""" + try: + model_info: Final = self.get_router_model_info( + deployment=cast(dict, deployment), # cast-ok: router deployments are plain dicts + received_model_name=model, + ) + except Exception as e: # noqa: BLE001 # best-effort: an unmappable deployment must not hide the others + verbose_router_logger.debug( + "litellm.router.py::_deployment_max_input_tokens: skipping deployment. Got - %s", e + ) + return None + max_input_tokens: Final = model_info.get("max_input_tokens") + return max_input_tokens if isinstance(max_input_tokens, int) else None + + def _pre_call_checks_need_token_count( + self, model: str, healthy_deployments: Sequence[Mapping[str, object]] + ) -> bool: + """Whether any healthy deployment declares a context window that a token count could exceed. + + Resolves each deployment the way ``_pre_call_checks`` does, so one unmappable deployment + cannot hide a later one that does declare a limit. + """ + return any( + self._deployment_max_input_tokens(model, deployment) is not None for deployment in healthy_deployments + ) + + async def _acount_pre_call_check_tokens( + self, + model: str, + healthy_deployments: Sequence[Mapping[str, object]], + messages: Sequence[Mapping[str, str]] | None, + input: str | Sequence[object] | None, + request_kwargs: Mapping[str, object] | None, + ) -> int | None: + """Count input tokens off the event loop, so a multi-MB prompt cannot stall the proxy. + + Returns None when no deployment limits its context window, and when counting fails. The + caller pairs this with ``skip_inline_token_count`` so neither case puts the count back on + the loop: a failed count leaves the deployments unfiltered, exactly as before. + """ + if messages is None and input is None: + return None + raw_instructions: Final = request_kwargs.get("instructions") if request_kwargs else None + try: + if not self._pre_call_checks_need_token_count(model, healthy_deployments): + return None + return await asyncify(self._count_pre_call_check_tokens)( + messages=cast(list[dict[str, str]] | None, messages), # cast-ok: forwarded to the sync counter + input=cast(str | list | None, input), # cast-ok: forwarded to the sync counter + instructions=raw_instructions if isinstance(raw_instructions, str) else None, + ) + except Exception as e: # noqa: BLE001 # best-effort: an uncountable prompt must not fail the request + verbose_router_logger.error( + "litellm.router.py::_acount_pre_call_check_tokens: failed to count tokens. Got - %s", e + ) + return None + def _pre_call_checks( self, model: str, @@ -10577,6 +10640,8 @@ class Router: messages: list[dict[str, str]] | None = None, input: str | list | None = None, request_kwargs: dict | None = None, + input_token_count: int | None = None, + skip_inline_token_count: bool = False, ): """ Filter out model in model group, if: @@ -10598,7 +10663,9 @@ class Router: # Token counting (tiktoken) is the dominant on-loop cost for large prompts. # Only count when a deployment actually declares max_input_tokens, and count # at most once; for model groups with no context-window limit it is skipped. - input_tokens: int | None = None + # Async callers pass the count in, already computed off the event loop, and set + # skip_inline_token_count so a failed off-loop count is not retried back on the loop. + input_tokens: int | None = input_token_count _context_window_error = False _potential_error_str = "" @@ -10633,6 +10700,8 @@ class Router: max_input_tokens = model_info.get("max_input_tokens") if isinstance(model_info, dict) else None if isinstance(max_input_tokens, int) and has_countable_input: if input_tokens is None: + if skip_inline_token_count: + return _returned_deployments try: input_tokens = self._count_pre_call_check_tokens( messages=messages, input=input, instructions=instructions @@ -11115,12 +11184,21 @@ class Router: ) if self.enable_pre_call_checks and (messages is not None or input is not None): + deployments_to_check: Final = cast(list[dict], healthy_deployments) healthy_deployments = self._pre_call_checks( model=model, - healthy_deployments=cast(list[dict], healthy_deployments), + healthy_deployments=deployments_to_check, messages=messages, input=input, request_kwargs=request_kwargs, + input_token_count=await self._acount_pre_call_check_tokens( + model=model, + healthy_deployments=deployments_to_check, + messages=messages, + input=input, + request_kwargs=request_kwargs, + ), + skip_inline_token_count=True, ) # check if user wants to do tag based routing healthy_deployments = await get_deployments_for_tag( diff --git a/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py b/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py index e928f4a0c3f..6e8406b2ec7 100644 --- a/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py +++ b/litellm/router_utils/pre_call_checks/prompt_caching_deployment_check.py @@ -9,6 +9,10 @@ from typing import Final, cast from litellm import verbose_logger from litellm.caching.dual_cache import DualCache from litellm.constants import DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT +from litellm.integrations.anthropic_cache_control_hook import ( + AllToolParamValues, + AnthropicCacheControlHook, +) from litellm.integrations.custom_logger import CustomLogger, Span from litellm.types.llms.openai import AllMessageValues from litellm.types.utils import CallTypes, StandardLoggingPayload @@ -63,8 +67,30 @@ class PromptCachingDeploymentCheck(CustomLogger): cache=self.cache, ) - model_id_dict: Final = await prompt_cache.async_get_model_id( + ## AUTO PROMPT CACHING - the breakpoints this request will carry are injected inside + ## `litellm.acompletion`, after a deployment has been picked, so the affinity key has to + ## be derived from the messages as they will be sent, not as they arrive here. + affinity_messages: Final = AnthropicCacheControlHook.messages_with_default_injections( messages=cast(list[AllMessageValues], messages), + models=( + deployment["litellm_params"]["model"] + for deployment in healthy_deployments + if isinstance(deployment.get("litellm_params"), dict) and deployment["litellm_params"].get("model") + ), + tools=( + cast( # cast-ok: request_kwargs is untyped; the stand-down scan duck-types every tool it reads + list[AllToolParamValues] | None, request_kwargs.get("tools") + ) + if request_kwargs is not None + else None + ), + enable_prompt_caching=( + request_kwargs.get("enable_prompt_caching") is True if request_kwargs is not None else None + ), + ) + + model_id_dict: Final = await prompt_cache.async_get_model_id( + messages=affinity_messages, tools=None, ) if model_id_dict is not None: diff --git a/litellm/rust_bridge/chat_completions.py b/litellm/rust_bridge/chat_completions.py new file mode 100644 index 00000000000..acda3086051 --- /dev/null +++ b/litellm/rust_bridge/chat_completions.py @@ -0,0 +1,453 @@ +"""Thin Python wrapper for the native Rust chat completions bridge. + +The Rust core owns the conversation translation, the provider call, and the +response normalization for the subset of `/chat/completions` requests it +accepts. This module only marshals inputs and hands the normalized result to +LiteLLM's existing `ModelResponse` builder. + +``None`` means the provider was never called, so the caller is free to serve the +request on the Python path. A failure after the call was issued raises instead: +retrying it there would bill the customer for the same work twice. +""" + +from __future__ import annotations + +import json +import os +from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass +from typing import TYPE_CHECKING, Final, Protocol + +import httpx +from pydantic import TypeAdapter, ValidationError + +from litellm._logging import verbose_logger +from litellm.exceptions import APIError +from litellm.litellm_core_utils.llm_response_utils.convert_dict_to_response import ( + convert_to_model_response_object, +) +from litellm.llms.bedrock.request_metadata import bedrock_request_metadata_is_owned +from litellm.rust_bridge.loader import get_native_bridge +from litellm.rust_bridge.timeouts import timeout_to_seconds +from litellm.types.utils import ModelResponse + +if TYPE_CHECKING: + from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + +# Providers whose `/chat/completions` deployments the Rust core can serve. A +# provider outside this set never reaches the bridge. +RUST_CHAT_COMPLETIONS_PROVIDERS: Final = frozenset({"anthropic", "bedrock"}) + +# `litellm_params` values are `object`, so validate the one this module reads +# rather than narrowing an unparameterized `Mapping` and typing the result Any. +_LITELLM_METADATA_ADAPTER: Final = TypeAdapter(Mapping[str, object]) + +RUST_RESPONSE_HEADER: Final = "x-litellm-rust" + +_TRUTHY_ENV_VALUES: Final = frozenset({"1", "true", "yes", "on"}) + + +class RustChatCompletions(Protocol): + def __call__( + self, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object] | None, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout_seconds: float | None, + ) -> Mapping[str, object]: + raise NotImplementedError + + +class RustAchatCompletions(Protocol): + def __call__( + self, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object] | None, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout_seconds: float | None, + ) -> Awaitable[Mapping[str, object]]: + raise NotImplementedError + + +class RustChatCompletionsDecline(Protocol): + def __call__( + self, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object] | None, + custom_llm_provider: str | None, + ) -> str | None: + raise NotImplementedError + + +class ResponseObserver(Protocol): + """Invoked with the payload the core returned, on success only. + + Lets the caller emit its own `post_call` on whichever path served the + request. Both entry points call it, so the synchronous and asynchronous + paths cannot drift apart the way the pre_call suppression once did. + """ + + def __call__(self, rust_response: Mapping[str, object], /) -> None: + raise NotImplementedError + + +def response_logger( + *, + logging_obj: LiteLLMLoggingObj, + messages: Sequence[object], + api_key: str, + additional_args: Mapping[str, object], +) -> ResponseObserver: + """A `ResponseObserver` that emits the caller's `post_call` for a Rust-served + request. + + The core owns the provider call, so the Python transform that normally + raises this event never runs; without it every `post_call` callback goes + silent on a Rust-served request and `original_response` stays unset. The + payload is the core's normalized response rather than the provider's wire + body, which is the closest thing that crosses the bridge. + """ + + def log(rust_response: Mapping[str, object], /) -> None: + logging_obj.post_call( + input=messages, + api_key=api_key, + original_response=json.dumps(rust_response), + additional_args=additional_args, + ) + + return log + + +class _Unset: + pass + + +_UNSET: Final[_Unset] = _Unset() + + +@dataclass(slots=True) +class _RustChatCompletionsState: + chat_completions: RustChatCompletions | None = None + achat_completions: RustAchatCompletions | None = None + decline: RustChatCompletionsDecline | None = None + + +_STATE: Final[_RustChatCompletionsState] = _RustChatCompletionsState() + + +def set_rust_chat_completions( + *, + chat_completions: RustChatCompletions | None | _Unset = _UNSET, + achat_completions: RustAchatCompletions | None | _Unset = _UNSET, + decline: RustChatCompletionsDecline | None | _Unset = _UNSET, +) -> None: + """Inject the native callables, so tests can supply a double instead of + patching module attributes.""" + if not isinstance(chat_completions, _Unset): + _STATE.chat_completions = chat_completions + if not isinstance(achat_completions, _Unset): + _STATE.achat_completions = achat_completions + if not isinstance(decline, _Unset): + _STATE.decline = decline + + +def load_rust_chat_completions() -> RustChatCompletions | None: + if _STATE.chat_completions is not None: + return _STATE.chat_completions + native_bridge: Final = get_native_bridge() + if native_bridge is None: + return None + loaded: RustChatCompletions | None = getattr(native_bridge, "chat_completions", None) + return loaded + + +def load_rust_achat_completions() -> RustAchatCompletions | None: + if _STATE.achat_completions is not None: + return _STATE.achat_completions + native_bridge: Final = get_native_bridge() + if native_bridge is None: + return None + loaded: RustAchatCompletions | None = getattr(native_bridge, "achat_completions", None) + return loaded + + +def _env_enables_rust() -> bool: + return os.getenv("LITELLM_RUST", "").strip().lower() in _TRUTHY_ENV_VALUES + + +def _load_rust_decline() -> RustChatCompletionsDecline | None: + if _STATE.decline is not None: + return _STATE.decline + native_bridge: Final = get_native_bridge() + if native_bridge is None: + return None + loaded: RustChatCompletionsDecline | None = getattr(native_bridge, "chat_completions_decline", None) + return loaded + + +def _anthropic_user_id_reaches_the_body(litellm_params: Mapping[str, object] | None) -> bool: + metadata: Final = litellm_params.get("metadata") if litellm_params is not None else None + try: + entries: Final = _LITELLM_METADATA_ADAPTER.validate_python(metadata) + except ValidationError: + return False + return entries.get("user_id") is not None + + +def _litellm_metadata_reaches_the_provider( + custom_llm_provider: str | None, litellm_params: Mapping[str, object] | None +) -> bool: + """Whether the Python transform would promote proxy-owned attribution into the + provider request, below this gate and inside the function the Rust route replaces. + + `AnthropicConfig.transform_request` promotes a valid `metadata["user_id"]` + into the Messages body, so the core never sees the key and would send the + request to Anthropic with the abuse-detection attribution missing. + + `AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the + Converse body whenever the operator armed `bedrock_request_metadata_fields`. + Owning that field also means evicting a caller-supplied one, which the core + cannot do either, so ownership alone is the condition rather than whether + anything resolved. + + Deliberately a superset of Python's condition in both cases: declining a + request Python would not have attributed anyway costs only the Rust path, + while missing one loses the attribution silently. + """ + match custom_llm_provider: + case "anthropic": + return _anthropic_user_id_reaches_the_body(litellm_params) + case "bedrock": + return bedrock_request_metadata_is_owned() + case _: + return False + + +def rust_chat_completions_accepts( + *, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object], + custom_llm_provider: str | None, + litellm_params: Mapping[str, object] | None, + stream: object, +) -> bool: + """Whether the Rust path will serve this request. + + Asked before the caller commits to either path, so pre-call logging is + emitted exactly once, on whichever path actually runs. The core's own + capability gate answers the second half; it resolves no credentials and + performs no I/O. + """ + if custom_llm_provider not in RUST_CHAT_COMPLETIONS_PROVIDERS: + return False + if stream: + return False + opted_in: Final = litellm_params is not None and litellm_params.get("rust") is True + if not opted_in and not _env_enables_rust(): + return False + if _litellm_metadata_reaches_the_provider(custom_llm_provider, litellm_params): + verbose_logger.debug("Rust chat completions declined (litellm metadata user_id); using the Python path") + return False + decline: Final = _load_rust_decline() + if decline is None: + return False + try: + reason: Final = decline( + model=model, + messages=messages, + optional_params=optional_params, + custom_llm_provider=custom_llm_provider, + ) + except Exception as rust_error: # noqa: BLE001 # rollout-safety fallback: any Rust bridge failure must fall back to the Python path + verbose_logger.debug( + "Rust chat completions gate raised %s; staying on the Python path", + type(rust_error).__name__, + ) + return False + if reason is not None: + verbose_logger.debug("Rust chat completions declined (%s); using the Python path", reason) + return False + return True + + +def _rust_bridge_exceptions() -> tuple[type[BaseException], type[BaseException]] | None: + """`(declined, upstream_failed)` from the native module, or None when absent.""" + native_bridge: Final = get_native_bridge() + if native_bridge is None: + return None + declined: Final = getattr(native_bridge, "RustBridgeDeclined", None) + upstream: Final = getattr(native_bridge, "RustUpstreamError", None) + if declined is None or upstream is None: + return None + return declined, upstream + + +def _reraise_or_decline( + rust_error: BaseException, + *, + model: str, + custom_llm_provider: str | None, +) -> None: + """Re-raise a failure the provider already saw, or return so the caller declines. + + A request that never reached the provider is safe to serve on the Python + path. One that did is not: the provider has already done the work, so a + second attempt bills for it twice. Those surface as an `APIError` carrying + the upstream status, which LiteLLM's exception mapping already understands. + """ + exceptions: Final = _rust_bridge_exceptions() + if exceptions is None: + verbose_logger.debug( + "Rust chat completions bridge raised %s; falling back to Python path", + type(rust_error).__name__, + ) + return + declined, upstream_failed = exceptions + if isinstance(rust_error, upstream_failed): + args: Final = rust_error.args + status: Final = args[0] if args else 0 + message: Final = args[1] if len(args) > 1 else "" + raise APIError( + status_code=int(status) or 500, + message=f"litellm rust chat completions: {message}", + llm_provider=custom_llm_provider or "", + model=model, + ) + if not isinstance(rust_error, declined): + raise rust_error + verbose_logger.debug( + "Rust chat completions declined before calling the provider (%s); using the Python path", + rust_error, + ) + + +def _build_model_response( + rust_response: Mapping[str, object], + model_response: ModelResponse, +) -> ModelResponse: + built: Final = convert_to_model_response_object( + response_object=dict(rust_response), # mutable-ok: the converter takes a real dict and rewrites it + model_response_object=model_response, + hidden_params={"additional_headers": {RUST_RESPONSE_HEADER: "true"}}, # mutable-ok: rewritten by the converter + ) + if not isinstance(built, ModelResponse): + raise TypeError(f"expected a ModelResponse from the rust path, got {type(built).__name__}") + return built + + +def chat_completions( + *, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object], + model_response: ModelResponse, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout: float | httpx.Timeout | None, + on_response: ResponseObserver, +) -> ModelResponse | None: + rust_chat_completions: Final = load_rust_chat_completions() + if rust_chat_completions is None: + return None + try: + rust_response: Final = rust_chat_completions( + model=model, + messages=messages, + optional_params=optional_params, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ) + except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw + _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) + return None + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + +async def achat_completions( + *, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object], + model_response: ModelResponse, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout: float | httpx.Timeout | None, + on_response: ResponseObserver, +) -> ModelResponse | None: + rust_achat_completions: Final = load_rust_achat_completions() + if rust_achat_completions is None: + return None + try: + rust_response: Final = await rust_achat_completions( + model=model, + messages=messages, + optional_params=optional_params, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout_seconds=timeout_to_seconds(timeout), + ) + except Exception as rust_error: # noqa: BLE001 # rollout safety: the helper re-raises anything the provider already saw + _reraise_or_decline(rust_error, model=model, custom_llm_provider=custom_llm_provider) + return None + on_response(rust_response) + return _build_model_response(rust_response, model_response) + + +async def achat_completions_or_fallback( + *, + model: str, + messages: Sequence[object], + optional_params: Mapping[str, object], + model_response: ModelResponse, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: Mapping[str, object] | None, + timeout: float | httpx.Timeout | None, + on_response: ResponseObserver, + python_fallback: Callable[[], Awaitable[object]], +) -> object: + """Await the Rust path, falling back to the caller's own Python path when + the bridge is unavailable or the call fails. + + The caller supplies the fallback, so the bridge stays free of provider + dispatch. This exists because a caller that dispatches asynchronously has + already returned a coroutine by the time a Rust failure surfaces, and so + cannot fall back on its own. + """ + response: Final = await achat_completions( + model=model, + messages=messages, + optional_params=optional_params, + model_response=model_response, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + timeout=timeout, + on_response=on_response, + ) + if response is not None: + return response + return await python_fallback() diff --git a/litellm/secret_managers/get_azure_ad_token_provider.py b/litellm/secret_managers/get_azure_ad_token_provider.py index d7f83855d2d..c2dc09bc65d 100644 --- a/litellm/secret_managers/get_azure_ad_token_provider.py +++ b/litellm/secret_managers/get_azure_ad_token_provider.py @@ -15,6 +15,8 @@ def infer_credential_type_from_environment() -> AzureCredentialType: and os.environ.get("AZURE_TENANT_ID") ): return AzureCredentialType.ClientSecretCredential + elif os.environ.get("AZURE_FEDERATED_TOKEN_FILE"): + return AzureCredentialType.DefaultAzureCredential elif os.environ.get("AZURE_CLIENT_ID"): return AzureCredentialType.ManagedIdentityCredential elif ( diff --git a/litellm/types/agents.py b/litellm/types/agents.py index 95562fcae8c..7e499dde642 100644 --- a/litellm/types/agents.py +++ b/litellm/types/agents.py @@ -1,7 +1,7 @@ from datetime import datetime from typing import TYPE_CHECKING, Any, Final, Literal -from pydantic import BaseModel, PrivateAttr +from pydantic import BaseModel, PrivateAttr, StrictInt from typing_extensions import Required, TypedDict from litellm.types.llms.base import LiteLLMPydanticObjectBase @@ -315,10 +315,23 @@ def _normalize_a2a_jsonrpc_response( The a2a SDK may omit ``id`` on error payloads even when the upstream agent returned it. Backfill from the outbound request id so LiteLLM can surface the agent error instead of failing Pydantic validation. + + JSON-RPC 2.0 requires the response id to equal the request id, so a string or + integer request id is carried over as-is. Anything else is stringified, which + is the only representation the response model accepts. + + A caller that supplied no id leaves the response id null, which is what the + spec requires for an error that cannot be correlated to a request. ``bool`` counts + as "anything else" despite subclassing ``int``, so ``true`` is never relayed as + ``1``, where it would collide with a real integer id. """ normalized: Final = dict(response_dict) - if normalized.get("id") is None and request_id is not None: - normalized["id"] = str(request_id) + if isinstance(normalized.get("id"), bool): + normalized["id"] = str(normalized["id"]) + elif normalized.get("id") is None and request_id is not None: + normalized["id"] = ( + request_id if isinstance(request_id, (str, int)) and not isinstance(request_id, bool) else str(request_id) + ) return normalized @@ -331,7 +344,7 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase): """ # A2A response fields - id: str + id: str | StrictInt | None = None jsonrpc: str = "2.0" result: dict[str, Any] | None = None error: dict[str, Any] | None = None @@ -360,8 +373,9 @@ class LiteLLMSendMessageResponse(LiteLLMPydanticObjectBase): Returns: LiteLLMSendMessageResponse with _hidden_params support """ - response_dict = response.model_dump(mode="json", exclude_none=True) - response_dict = _normalize_a2a_jsonrpc_response(response_dict, request_id=request_id) + response_dict: Final = _normalize_a2a_jsonrpc_response( + response.model_dump(mode="json", exclude_none=True), request_id=request_id + ) return cls(**response_dict) @classmethod diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index d68d1dc9625..e2469d4c78f 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -169,6 +169,10 @@ ShadowEvalDirection: TypeAlias = Literal["forward", "reverse"] DEFAULT_SHADOW_EVAL_JUDGE_MODEL: Final[str] = "anthropic/claude-sonnet-5" +# Sample-count ceiling written on every new job: a zero-cost error loop (a shadow arm that +# fails before billing) never consumes spend budget, so it must terminate on count instead. +SHADOW_EVAL_TURN_VALVE: Final[int] = 10_000 + class StartShadowEvalRequest(BaseModel): """Start duplicating one or more keys' traffic for blind comparison against an auto-router.""" @@ -179,7 +183,7 @@ class StartShadowEvalRequest(BaseModel): description=( "The hashed virtual keys whose traffic will be shadowed. Shadow evaluation runs ONLY on these " "keys' traffic; requests made with any other key are not sampled. Each key carries its own " - "max_turns budget, so one key exhausting its budget leaves the others sampling. At most 100 " + "max_budget spend budget, so one key exhausting its budget leaves the others sampling. At most 100 " "keys per job, which also bounds every read the job's endpoints make." ), ) @@ -219,17 +223,27 @@ class StartShadowEvalRequest(BaseModel): le=30, description="How many days the job samples traffic before completing on its own", ) - max_turns: int = Field( - default=200, - ge=1, - le=2000, + max_budget: float = Field( + default=10.0, + ge=0.01, + le=10_000, description=( - "Per-key sample budget: the job judges at most this many turns of EACH scoped key's traffic, " - "so a job over N keys judges at most N times max_turns turns. This is also the spend bound; " - "expected judge cost is roughly that turn ceiling times one judge call" + "Per-key USD budget for the eval's own overhead, the shadow-arm and judge calls, priced with " + "the same figures the spend pipeline bills. EACH scoped key samples until its recorded eval " + "spend reaches this, so a job over N keys spends at most about N times max_budget; in-flight " + "samples can overshoot the cap by one sampling cache window" ), ) + @model_validator(mode="before") + @classmethod + def _reject_the_retired_turn_budget(cls, values: object) -> object: + """Pydantic ignores unknown fields, so a caller still sending max_turns would + silently run on the default dollar budget instead of the bound they asked for.""" + if isinstance(values, Mapping) and "max_turns" in values: + raise ValueError("max_turns was replaced by max_budget, the per-key USD cap on the eval's own spend") + return values + @field_validator("shadow_percentage") @classmethod def _round_percentage(cls, value: float) -> float: @@ -296,7 +310,19 @@ class ShadowEvalJobKeyResponse(BaseModel): """One key a job shadows, with its own budget and stop state.""" api_key_id: str = Field(description="The hashed virtual key whose traffic this entry scopes") - max_turns: int = Field(description="This key's own sample budget, independent of its siblings'") + max_turns: int = Field( + description=( + "This key's sample-count ceiling: the whole budget for jobs created before max_budget " + "existed, and the error-loop safety valve otherwise" + ) + ) + max_budget: float | None = Field( + default=None, + description=( + "This key's own USD budget for the eval's shadow and judge spend, independent of its " + "siblings'; None on jobs created before spend budgets existed, which max_turns alone bounds" + ), + ) stopped_at: datetime | None = Field( default=None, description=( @@ -313,10 +339,19 @@ class ShadowEvalJobKeyResponse(BaseModel): "once the key is stamped, so in-flight attempts landing after a stop never reclassify it" ), ) + spend: float | None = Field( + default=None, + description=( + "This key's recorded shadow plus judge spend in USD, the same figure the sampler budgets " + "against max_budget; populated on list and detail responses and frozen at stopped_at " + "exactly like attempt_count" + ), + ) @property def budget_spent(self) -> bool: - return self.attempt_count is not None and self.attempt_count >= self.max_turns + over_spend: Final = self.max_budget is not None and self.spend is not None and self.spend >= self.max_budget + return over_spend or (self.attempt_count is not None and self.attempt_count >= self.max_turns) key_alias: str | None = Field( default=None, diff --git a/litellm/utils.py b/litellm/utils.py index 867f7a93452..b8a9ea37e05 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -5245,7 +5245,7 @@ def _check_provider_match(model_info: dict, custom_llm_provider: str | None) -> return True -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict class PotentialModelNamesAndCustomLLMProvider(TypedDict): @@ -5253,6 +5253,7 @@ class PotentialModelNamesAndCustomLLMProvider(TypedDict): combined_model_name: str stripped_model_name: str combined_stripped_model_name: str + provider_prefixed_model_name: ReadOnly[str] custom_llm_provider: str @@ -5280,6 +5281,7 @@ def _get_model_info_from_generalization( potential_model_names["split_model"], potential_model_names["combined_stripped_model_name"], potential_model_names["stripped_model_name"], + potential_model_names["provider_prefixed_model_name"], ) if any(_get_model_cost_key(candidate) is not None for candidate in candidates): return None @@ -5304,6 +5306,7 @@ def _get_potential_model_names(model: str, custom_llm_provider: str | None) -> P combined_model_name = model stripped_model_name = _strip_model_name(model=model, custom_llm_provider=custom_llm_provider) combined_stripped_model_name = stripped_model_name + provider_prefixed_model_name = model elif custom_llm_provider and model.startswith( custom_llm_provider + "/" ): # handle case where custom_llm_provider is provided and model starts with custom_llm_provider @@ -5311,11 +5314,13 @@ def _get_potential_model_names(model: str, custom_llm_provider: str | None) -> P combined_model_name = model stripped_model_name = _strip_model_name(model=split_model, custom_llm_provider=custom_llm_provider) combined_stripped_model_name = f"{custom_llm_provider}/{stripped_model_name}" + provider_prefixed_model_name = f"{custom_llm_provider}/{model}" else: split_model = model combined_model_name = f"{custom_llm_provider}/{model}" stripped_model_name = _strip_model_name(model=model, custom_llm_provider=custom_llm_provider) combined_stripped_model_name = f"{custom_llm_provider}/{stripped_model_name}" + provider_prefixed_model_name = combined_model_name if custom_llm_provider in ("bedrock", "bedrock_converse"): from litellm.llms.bedrock.common_utils import strip_bedrock_routing_prefix @@ -5327,6 +5332,7 @@ def _get_potential_model_names(model: str, custom_llm_provider: str | None) -> P combined_model_name=combined_model_name, stripped_model_name=stripped_model_name, combined_stripped_model_name=combined_stripped_model_name, + provider_prefixed_model_name=provider_prefixed_model_name, custom_llm_provider=cast(str, custom_llm_provider), ) @@ -5435,6 +5441,7 @@ def _get_model_info_helper( combined_model_name: Final = potential_model_names["combined_model_name"] stripped_model_name: Final = potential_model_names["stripped_model_name"] combined_stripped_model_name: Final = potential_model_names["combined_stripped_model_name"] + provider_prefixed_model_name: Final = potential_model_names["provider_prefixed_model_name"] split_model: Final = potential_model_names["split_model"] custom_llm_provider = potential_model_names["custom_llm_provider"] model_cost_custom_llm_provider: Final = custom_llm_provider @@ -5493,6 +5500,10 @@ def _get_model_info_helper( 3. 'split_model' in litellm.model_cost. Checks "au.anthropic.claude-opus-4-8" in litellm.model_cost if model="bedrock/au.anthropic.claude-opus-4-8" 4. 'combined_stripped_model_name' in litellm.model_cost. Checks if 'gemini/gemini-1.5-flash' in model map, if 'gemini/gemini-1.5-flash-001' given. 5. 'stripped_model_name' in litellm.model_cost. Checks if 'ft:gpt-3.5-turbo' in model map, if 'ft:gpt-3.5-turbo:my-org:custom_suffix:id' given. + 6. 'provider_prefixed_model_name' in litellm.model_cost, for providers whose own model ids repeat the + litellm provider name. Checks "perplexity/perplexity/glm-5.2" if model="perplexity/glm-5.2" and + custom_llm_provider="perplexity", where 1-5 all read the leading "perplexity/" as the litellm prefix + and strip it. Tried last so no model that already resolves through 1-5 can change. """ _model_info: dict[str, Any] | None = None @@ -5548,6 +5559,16 @@ def _get_model_info_helper( custom_llm_provider=model_cost_custom_llm_provider, ): _model_info = None + if _model_info is None: + _matched_key = _get_model_cost_key(provider_prefixed_model_name) + if _matched_key is not None: + key = _matched_key + _model_info = _get_model_info_from_model_cost(key=cast(str, key)) + if not _check_provider_match( + model_info=_model_info, + custom_llm_provider=model_cost_custom_llm_provider, + ): + _model_info = None if _model_info is None: generalization: Final = _get_model_info_from_generalization( diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b9c8824aa67..858eab672e5 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -34758,6 +34758,50 @@ "supports_reasoning": false, "supports_function_calling": true }, + "perplexity/perplexity/deepseek-v4-flash-0731": { + "cache_read_input_token_cost": 2.8e-08, + "input_cost_per_token": 1.3e-07, + "litellm_provider": "perplexity", + "mode": "responses", + "output_cost_per_token": 2.6e-07, + "source": "https://docs.perplexity.ai/docs/agent-api/models", + "supports_web_search": true, + "supports_reasoning": true, + "supports_function_calling": true + }, + "perplexity/perplexity/glm-5.2": { + "cache_read_input_token_cost": 1.4e-07, + "input_cost_per_token": 1.4e-06, + "litellm_provider": "perplexity", + "mode": "responses", + "output_cost_per_token": 4.4e-06, + "source": "https://docs.perplexity.ai/docs/agent-api/models", + "supports_web_search": true, + "supports_reasoning": true, + "supports_function_calling": true + }, + "perplexity/perplexity/kimi-k3": { + "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 3e-06, + "litellm_provider": "perplexity", + "mode": "responses", + "output_cost_per_token": 1.5e-05, + "source": "https://docs.perplexity.ai/docs/agent-api/models", + "supports_web_search": true, + "supports_reasoning": true, + "supports_function_calling": true + }, + "perplexity/perplexity/kimi-k2.7-code": { + "cache_read_input_token_cost": 1.9e-07, + "input_cost_per_token": 9.5e-07, + "litellm_provider": "perplexity", + "mode": "responses", + "output_cost_per_token": 4e-06, + "source": "https://docs.perplexity.ai/docs/agent-api/models", + "supports_web_search": true, + "supports_reasoning": false, + "supports_function_calling": true + }, "perplexity/pplx-embed-v1-0.6b": { "input_cost_per_token": 4e-09, "litellm_provider": "perplexity", diff --git a/pyproject.toml b/pyproject.toml index 09a69f3771e..6e3c181ae1d 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -67,8 +67,8 @@ proxy = [ "azure-identity>=1.25.2,<2.0", "azure-storage-blob>=12.28.0,<13.0", "mcp>=1.28.1,<2.0", - "litellm-proxy-extras==0.4.87", - "litellm-enterprise==0.1.57", + "litellm-proxy-extras==0.4.88", + "litellm-enterprise==0.1.58", "RestrictedPython>=8.1,<9.0", "rich>=13.9.4,<14.0", "InquirerPy>=0.3.4,<1.0", diff --git a/ruff-strict-budget.json b/ruff-strict-budget.json index 5c312dcf1c8..a990f7c3830 100644 --- a/ruff-strict-budget.json +++ b/ruff-strict-budget.json @@ -57,7 +57,7 @@ "limit": 3 }, "BLE001": { - "limit": 2923 + "limit": 2920 }, "C401": { "limit": 8 @@ -174,7 +174,7 @@ "limit": 176 }, "RUF012": { - "limit": 241 + "limit": 240 }, "RUF015": { "limit": 8 diff --git a/ruff-tests.toml b/ruff-tests.toml new file mode 100644 index 00000000000..c760a53ed91 --- /dev/null +++ b/ruff-tests.toml @@ -0,0 +1,24 @@ +# Lint config for the test tree, which ruff.toml excludes from `ruff check`. +# +# Every rule here catches a test that cannot fail. Rules land one at a time, each +# with its existing violations already fixed, so this list never needs a budget +# file or a ratchet. +# +# F821 a name that does not exist raises NameError, and a test body wrapped in +# `except Exception: pass` swallows that NameError and reports green +# B011 `assert False` inside `try:` raises AssertionError, which the `except +# Exception` below it catches. `pytest.fail` raises BaseException and escapes +# PT015 same site as B011, from the pytest ruleset +# B015 a bare `a == b` statement is evaluated and thrown away; the missing `assert` +# means the test checks nothing +# B018 a bare attribute access or literal, usually a call missing its parens +# PLW0127 `x = x` self-assignment, dead code that reads like a narrowing or a fixup +# PLR0133 comparison of two constants, e.g. `assert True == True` +# +# No target-version here on purpose: it resolves from requires-python (>=3.10), so +# 3.11-only builtins like BaseExceptionGroup are correctly flagged in a tree that +# still has to run on 3.10. + +line-length = 120 + +lint.select = ["F821", "B011", "B015", "B018", "PT015", "PLR0133", "PLW0127"] diff --git a/schema.prisma b/schema.prisma index 60058c777ca..d9959677116 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1502,7 +1502,8 @@ model LiteLLM_ShadowEvalJob { baseline_model String? // reverse only: the fixed model the router is judged against judge_model String shadow_percentage Float - max_turns Int // this key's sample budget: judge at most this many turns + max_turns Int // sample-count ceiling: the whole budget on pre-max_budget jobs, the error-loop valve otherwise + max_budget Float? // per-key USD cap on the eval's own shadow + judge spend; null on jobs from before spend budgets created_at DateTime @default(now()) created_by String? ends_at DateTime @@ -1525,6 +1526,7 @@ model LiteLLM_ShadowEvalAttempt { shadow_model String? confidence Float? judge_cost Float @default(0) + shadow_cost Float @default(0) error String? created_at DateTime @default(now()) diff --git a/scripts/check_test_quality.py b/scripts/check_test_quality.py index 5b0b03c60fb..e3ffbac9808 100644 --- a/scripts/check_test_quality.py +++ b/scripts/check_test_quality.py @@ -39,15 +39,23 @@ TQ005 `litellm. = ...` module-global mutation. The SDK's module globals what the 491-line save/restore conftest exists to paper over. Inject the dependency or use a fixture that restores it. TQ006 A `pytest.skip` reached only when a credential-shaped environment variable is - absent. Absence is what the condition has to say: `not key`, `key is None`, - `"KEY" not in os.environ`. A skip taken when the credential is present is - somebody's deliberate branch and is left alone. On a runner that does not hold that credential the guard fires every + absent. On a runner that does not hold that credential the guard fires every time, so the test reports green having executed nothing and is indistinguishable from coverage that exists. Fake the provider at the HTTP boundary, or fail - loudly, so a missing credential shows up as a missing credential. The gate is - followed through one local or module-level binding, which is the - `key = os.getenv(...)` then `if not key: pytest.skip(...)` shape most of these - use. + loudly, so a missing credential shows up as a missing credential. Absence is + what the condition has to say -- `not key`, `key is None`, `"KEY" not in + os.environ` -- since a skip taken when the credential is present is somebody's + deliberate branch. The gate follows one local or module-level binding, which is + the `key = os.getenv(...)` then `if not key: pytest.skip(...)` shape most of + these use. +TQ007 A module global that a conftest saves before every test and restores after it. + The save/restore list is a hand-maintained inventory of the leaks the suite + already knows about, so it is allowed to shrink and never to grow: a new entry + means one more global whose lifetime the tests manage instead of the code owning + it. Give the consumers an injection seam rather than another snapshot line. The + names are read from the keys the conftest assigns directly and from whatever the + save loop iterates, including a module-level tuple or dict it names rather than + spells out. Every rule is suppressible with `# test-quality-ok: ` on the reported line, following the repo's `*-ok: ` convention. A suppression without a @@ -123,6 +131,9 @@ PATCH_MEMBERS: Final = frozenset(("object", "dict", "multiple")) ENVIRON_READERS: Final = frozenset(("os.environ.get", "environ.get", "os.getenv", "getenv")) ENVIRON_MAPPINGS: Final = frozenset(("os.environ", "environ")) SKIP_CALLS: Final = frozenset(("pytest.skip", "skip")) +CONFTEST_NAME: Final = "conftest.py" +SDK_MODULE: Final = "litellm" + CREDENTIAL_NAME_RE: Final = re.compile( r"(?:API_KEY|_KEY|TOKEN|SECRET|PASSWORD|CREDENTIAL|DATABASE_URL|ACCESS_KEY_ID)$" ) @@ -547,6 +558,105 @@ def iter_credential_skip_violations(path: Path, tree: ast.Module) -> Iterator[Vi ) +def _reads_sdk_attribute(node: ast.AST) -> bool: + return any( + ( + isinstance(inner, ast.Call) + and _dotted_name(inner.func) == "getattr" + and bool(inner.args) + and _dotted_name(inner.args[0]) == SDK_MODULE + ) + or (isinstance(inner, ast.Attribute) and _dotted_name(inner.value) == SDK_MODULE) + for inner in ast.walk(node) + ) + + +def _subscript_targets(node: ast.AST) -> Iterator[ast.Subscript]: + for inner in ast.walk(node): + if isinstance(inner, ast.Assign): + yield from (target for target in inner.targets if isinstance(target, ast.Subscript)) + + +def _saves_sdk_attribute_by_key(node: ast.AST) -> Iterator[ast.Subscript]: + """Every `["name"] = `, whatever the dict is called. + + Matching on the shape rather than on a list of blessed dict names is what reaches + the conftest that builds its snapshot inside a helper and calls the dict `state`. + """ + for inner in ast.walk(node): + if isinstance(inner, ast.Assign) and _reads_sdk_attribute(inner.value): + yield from (target for target in inner.targets if isinstance(target, ast.Subscript)) + + +def _saves_sdk_attributes_in_loop(node: ast.For) -> bool: + """A save loop reads the SDK and stores under the loop variable, in either order. + + The read is often bound to a local first (`val = getattr(litellm, attr)`) and only + then stored, so the read and the store are separate statements and cannot be + required of the same assignment. + """ + if not isinstance(node.target, ast.Name): + return False + stores_by_key: Final = any( + isinstance(subscript.slice, ast.Name) and subscript.slice.id == node.target.id + for statement in node.body + for subscript in _subscript_targets(statement) + ) + return stores_by_key and any(_reads_sdk_attribute(statement) for statement in node.body) + + +def _module_constants(tree: ast.Module) -> Mapping[str, ast.expr]: + return MappingProxyType({ + target.id: node.value + for node in tree.body + if isinstance(node, ast.Assign) + for target in node.targets + if isinstance(target, ast.Name) + }) + + +def _string_members(node: ast.expr) -> Iterator[tuple[str, int]]: + """The string names a collection literal holds: a tuple/list's items, a dict's keys.""" + elements: Final = ( + node.elts if isinstance(node, (ast.Tuple, ast.List)) else node.keys if isinstance(node, ast.Dict) else () + ) + yield from ( + (element.value, element.lineno) + for element in elements + if isinstance(element, ast.Constant) and isinstance(element.value, str) + ) + + +def _snapshotted_names(tree: ast.Module) -> Iterator[tuple[str, int]]: + constants: Final = _module_constants(tree) + for node in ast.walk(tree): + if isinstance(node, ast.Assign): + yield from ( + (subscript.slice.value, subscript.lineno) + for subscript in _saves_sdk_attribute_by_key(node) + if isinstance(subscript.slice, ast.Constant) and isinstance(subscript.slice.value, str) + ) + elif isinstance(node, ast.For) and _saves_sdk_attributes_in_loop(node): + iterable: Final = constants.get(node.iter.id) if isinstance(node.iter, ast.Name) else node.iter + if iterable is not None: + yield from _string_members(iterable) + + +def iter_conftest_inventory_violations(path: Path, tree: ast.Module) -> Iterator[Violation]: + if path.name != CONFTEST_NAME: + return + seen: Final = dict(reversed(tuple(_snapshotted_names(tree)))) + for name, line in sorted(seen.items(), key=lambda item: item[1]): + yield Violation( + path, + line, + "TQ007", + f"`litellm.{name}` is saved and restored around every test in this tree; the list is an " + "inventory of known leaks and may only shrink, so give the consumers an injection seam " + f"instead of adding to it (suppress: `# {SUPPRESSION_TOKEN}: `)", + ) + + def check_file(path: Path) -> tuple[Violation, ...]: try: source: Final = path.read_text(encoding="utf-8") @@ -567,6 +677,7 @@ def check_file(path: Path) -> tuple[Violation, ...]: *iter_environ_violations(path, tree), *iter_global_mutation_violations(path, tree), *iter_credential_skip_violations(path, tree), + *iter_conftest_inventory_violations(path, tree), ) if violation.line not in skip ) diff --git a/scripts/test_quality_gate.py b/scripts/test_quality_gate.py index 292da29e1f2..7d34b194f1c 100644 --- a/scripts/test_quality_gate.py +++ b/scripts/test_quality_gate.py @@ -12,9 +12,15 @@ Every rule is seeded at exactly its count on the day the gate landed, so the suite's existing debt is grandfathered and any net-new violation trips the gate immediately. ``--update`` ratchets a limit down by the violations this branch fixed relative to its branch point (the merge-base), so the ceilings only ever -fall. A rule absent from the budget at the merge-base was seeded on this branch; -``--update`` leaves its limit untouched, because the base tree predates the rule -and its whole grandfathered count would otherwise be misread as "fixed". +fall. Base counts are measured with the *current* checker, so a rule introduced +on this branch is counted at the base too and ratchets like every other one. + +Only ever falling is not the same as always falling, so the gate enforces the +second half: a branch that clears violations and leaves the ceiling above its +new count fails, naming the rules and telling the author to run +``make lint-budget-update``. Without that, a removed violation could come back +later under a ceiling nobody lowered. Drift already in the base is never +blamed, so this fires only on the branch that did the clearing. The deliberate difference from its sibling: this gate has no headroom anywhere. Type discipline seeded LIT010/LIT011 at 1.5x to leave room for an in-flight @@ -138,6 +144,21 @@ def over_ceiling(head: Mapping[str, int], budget: Mapping[str, Mapping[str, int] ) +def unratcheted( + head: Mapping[str, int], + base: Mapping[str, int], + budget: Mapping[str, Mapping[str, int]], +) -> tuple[Breach, ...]: + """Rules this branch cleared without lowering the ceiling behind them. Requires + both `head < base`, so drift already in the base is never blamed on this change, + and `head < limit`, so a ceiling already at the count is left alone.""" + return tuple(sorted( + Breach(rule, head.get(rule, 0), spec["limit"], head.get(rule, 0) - base.get(rule, 0)) + for rule, spec in budget.items() + if head.get(rule, 0) < base.get(rule, 0) and head.get(rule, 0) < spec["limit"] + )) + + def evaluate( head: Mapping[str, int], base: Mapping[str, int], @@ -177,15 +198,39 @@ def introduced( return tuple(v for v in violations if v.line in changed.get(v.file, frozenset())) +def touches_measured_tree(base_point: str) -> bool: + """Whether this branch changed anything that can move a count. A branch that + touches neither the test tree nor the checker cannot have cleared a violation, + so the base scan is skipped and the gate stays cheap on the common change.""" + changed: Final = _run( + ["git", "diff", "--name-only", base_point, "--", TARGET, str(CHECKER.relative_to(REPO_ROOT))] + ) + return bool(changed.strip()) + + def cmd_check(base: str) -> None: budget: Final = json.loads(BUDGET_PATH.read_text()) head: Final = head_violations() head_counts: Final = count_by_rule(head) - if not over_ceiling(head_counts, budget): + base_point: Final = resolve_base_point(base) + if not over_ceiling(head_counts, budget) and not touches_measured_tree(base_point): print(f"OK: every TQ rule is within its test-suite ceiling (base {base})") return - base_point: Final = resolve_base_point(base) - breaches: Final = evaluate(head_counts, base_counts(base_point), budget) + base_at_point: Final = base_counts(base_point) + stale: Final = unratcheted(head_counts, base_at_point, budget) + if stale: + print(f"FAIL: TQ-rule limits were left above the count this branch reached (base {base}):") + for breach in stale: + print( + f" {breach.rule}: this branch cleared {-breach.added} down to {breach.total}, " + f"but the limit is still {breach.cap}" + ) + print( + "Run `make lint-budget-update` and commit the lowered limits, so the " + "violations you cleared cannot come back under a ceiling nobody moved." + ) + raise SystemExit(1) + breaches: Final = evaluate(head_counts, base_at_point, budget) if not breaches: print(f"OK: every TQ rule is within its test-suite ceiling (base {base})") return @@ -216,46 +261,25 @@ def ratcheted_budget( budget: Mapping[str, Mapping[str, int]], current: Mapping[str, int], base: Mapping[str, int], - seeded: frozenset[str] = frozenset(), ) -> Mapping[str, Mapping[str, int]]: """Each rule's limit lowered by the violations `current` fixed vs `base`. The drop - is clamped to what was actually cleared, so a limit only ever falls. Rules in - `seeded` were introduced on this branch and pass through untouched.""" + is clamped to what was actually cleared, so a limit only ever falls.""" return MappingProxyType({ - rule: { - "limit": spec["limit"] if rule in seeded - else max(0, spec["limit"] - max(0, base.get(rule, 0) - current.get(rule, 0))) - } + rule: {"limit": max(0, spec["limit"] - max(0, base.get(rule, 0) - current.get(rule, 0)))} for rule, spec in sorted(budget.items()) }) -def _base_budget_rules(base_point: str) -> frozenset[str]: - proc: Final = subprocess.run( - ["git", "show", f"{base_point}:{BUDGET_PATH.name}"], - cwd=REPO_ROOT, capture_output=True, text=True, - ) - if proc.returncode != 0: - return frozenset() - return frozenset(json.loads(proc.stdout)) - - def cmd_update(base_ref: str = DEFAULT_BASE) -> None: """Ratchet each rule's limit down by the violations this branch fixed.""" budget: Final = json.loads(BUDGET_PATH.read_text()) base_point: Final = resolve_base_point(base_ref) - seeded: Final = frozenset(budget) - _base_budget_rules(base_point) updated: Final = ratcheted_budget( - budget, count_by_rule(head_violations()), base_counts(base_point), seeded + budget, count_by_rule(head_violations()), base_counts(base_point) ) BUDGET_PATH.write_text(json.dumps(dict(updated), indent=2, sort_keys=True) + "\n") cleared: Final = sum(budget[rule]["limit"] - updated[rule]["limit"] for rule in updated) print(f"Ratcheted TQ-rule limits down by {cleared} violations this branch fixed") - if seeded: - print( - "Left untouched (seeded on this branch, absent from the base budget): " - + ", ".join(sorted(seeded)) - ) def cmd_seed() -> None: diff --git a/test-quality-budget.json b/test-quality-budget.json index 2a5945fe36c..17baf64601b 100644 --- a/test-quality-budget.json +++ b/test-quality-budget.json @@ -16,5 +16,8 @@ }, "TQ006": { "limit": 34 + }, + "TQ007": { + "limit": 117 } } diff --git a/tests/_fake_openai_endpoint_server.py b/tests/_fake_openai_endpoint_server.py index caf3fb5ba2a..ac83e74b66a 100644 --- a/tests/_fake_openai_endpoint_server.py +++ b/tests/_fake_openai_endpoint_server.py @@ -8,11 +8,11 @@ those jobs failed with ``404 Application not found`` even though nothing in the PR was broken. This process is the local stand-in. A model points its ``api_base`` here and -gets back a well-formed chat/text/embedding response with realistic ``usage`` so -cost tracking and spend accounting still exercise their real code paths. The one -behavioral special case mirrors the old hosted mock: a request whose ``model`` -is ``429`` returns HTTP 429 so rate-limit and cooldown tests still have -something to trip on. +gets back a well-formed chat/text/embedding/moderation response with realistic +``usage`` so cost tracking and spend accounting still exercise their real code +paths. The one behavioral special case mirrors the old hosted mock: a request +whose ``model`` is ``429`` returns HTTP 429 so rate-limit and cooldown tests +still have something to trip on. """ from __future__ import annotations @@ -35,6 +35,21 @@ _SLOW_MODEL: Final = "slow-endpoint" _SLOW_RESPONSE_SECONDS: Final = 3.0 _PROMPT_TOKENS: Final = 20 _COMPLETION_TOKENS: Final = 20 +_MODERATION_CATEGORIES: Final = ( + "harassment", + "harassment/threatening", + "hate", + "hate/threatening", + "illicit", + "illicit/violent", + "self-harm", + "self-harm/instructions", + "self-harm/intent", + "sexual", + "sexual/minors", + "violence", + "violence/graphic", +) def _usage() -> dict[str, int]: @@ -220,6 +235,28 @@ async def triton_embeddings(_request: Request) -> Response: ) +def _moderation_result() -> dict[str, object]: + return { + "flagged": False, + "categories": {category: False for category in _MODERATION_CATEGORIES}, + "category_scores": {category: 0.0 for category in _MODERATION_CATEGORIES}, + "category_applied_input_types": {category: ["text"] for category in _MODERATION_CATEGORIES}, + } + + +async def moderations(request: Request) -> Response: + body: Final = await _parse_body(request) + raw_input: Final = body.get("input", "") + count: Final = len(raw_input) if isinstance(raw_input, list) else 1 + return JSONResponse( + { + "id": f"modr-{uuid.uuid4().hex[:24]}", + "model": _requested_model(body), + "results": [_moderation_result() for _ in range(max(count, 1))], + } + ) + + async def list_models(_request: Request) -> Response: return JSONResponse( { @@ -247,6 +284,8 @@ app = Starlette( Route("/embeddings", embeddings, methods=["POST"]), Route("/v1/embeddings", embeddings, methods=["POST"]), Route("/triton/embeddings", triton_embeddings, methods=["POST"]), + Route("/moderations", moderations, methods=["POST"]), + Route("/v1/moderations", moderations, methods=["POST"]), Route("/models", list_models, methods=["GET"]), Route("/v1/models", list_models, methods=["GET"]), ] diff --git a/tests/documentation_tests/test_router_settings.py b/tests/documentation_tests/test_router_settings.py index 290aa283af4..a1b6f1dac1d 100644 --- a/tests/documentation_tests/test_router_settings.py +++ b/tests/documentation_tests/test_router_settings.py @@ -61,7 +61,7 @@ try: documented_keys.update(doc_key_pattern.findall(table_content)) except Exception as e: raise Exception( - f"Error reading documentation: {e}, \n repo base - {os.listdir(repo_base)}" + f"Error reading documentation: {e}, \n repo base - {os.listdir(_repo_root)}" ) diff --git a/tests/e2e/CLAUDE.md b/tests/e2e/CLAUDE.md index d7334552d0c..840a40a54cd 100644 --- a/tests/e2e/CLAUDE.md +++ b/tests/e2e/CLAUDE.md @@ -73,13 +73,17 @@ Mark live tests with `@pytest.mark.e2e` (on the class or the module). Pure cover ## Record and replay fixtures -`E2E_FIXTURE_MODE` selects the transport every client is built on: `live` (the default, and what an unset variable means: nothing changes), `record` (run against the live proxy and write every interaction to a fixture bundle), or `replay` (serve every interaction back from the bundle with no HTTP at all, so a replay run needs no proxy and cannot bill a provider). The seam is `select_transport` in `fixture_transport.py`, applied inside `build_proxy_client`; both transports fulfil the same `Transport` protocol, so no test or client changes shape in any mode +`E2E_FIXTURE_MODE` scopes the proxy's provider-bound traffic: `live` (the default, and what an unset variable means: nothing changes), `record` (the proxy's provider calls are forwarded to the real provider through a local edge server and written to a fixture bundle), or `replay` (the edge answers those calls from the bundle, so the run makes zero provider calls and spends nothing). Test-to-proxy traffic always goes over the wire in every mode: record and replay both need the live proxy and database, because the point is that key auth, routing, cost calculation, and spend-log writes execute for real while only the provider is swapped out. Breaking any of those in the proxy turns a replay run red -A bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`) is a directory: `manifest.json` carries the record timestamp, harness git version, and format version, and each test gets a subdirectory holding one JSON file per transport call in call order (`0000-post-chat-completions.json`). Auth header values and credential request fields (`api_key`, `*_secret_key`, `static_headers`, and the like; the list is `fixture_canonical.py`'s) are redacted on write, and file uploads store a sha256 digest instead of the bytes; response bodies are stored verbatim (a /key/generate response keeps the ephemeral virtual key it minted), which is part of why bundles are gitignored. `fixture_bundle.py` owns the format +The seam is `provider_edge.py`: `start_provider_edge` boots an in-process HTTP server (one shared instance per pytest process, `e2e_config.provider_edge_base` is the accessor) that mounts each supported provider under a path prefix (`EDGE_MOUNTS`: `/openai` -> `https://api.openai.com`, `/anthropic` -> `https://api.anthropic.com`). A test participates by registering its deployment with `api_base=provider_edge_base("openai")` plus the provider's path suffix; `quota_management/spend_tracking/test_provider_edge_spend_e2e.py` is the reference. In live mode the accessor returns None and the deployment defaults to the real provider, so an edge-wired test runs in all three modes unchanged. Non-wired tests hit their providers live in every mode. The edge binds `E2E_PROVIDER_EDGE_BIND_HOST` (default 127.0.0.1) and advertises `E2E_PROVIDER_EDGE_ADVERTISE_HOST` in the api_base it hands out, for proxies running in containers -Replay matches calls per test by canonical key: `fixture_canonical.py` canonicalizes the recorded request (volatile headers and credential fields out, unique markers, generated ids, uuids, and timestamps replaced with fixed placeholders, object keys sorted) and the key is the method, path, and a content hash, so identity survives re-records and machine changes while any real content drift is a `ReplayMiss` that names the computed key, the closest recorded key with its file, and a content diff, and never falls through to a live call. Matching is order-independent across distinct keys (concurrent calls may interleave) and FIFO within one key (a poll loop replays its responses in recorded order); a passed test must also consume its whole recording, or teardown fails it naming a leftover key. Either way the fix is always to re-record with `E2E_FIXTURE_MODE=record`. Every rewrite rule lives in `fixture_canonical.py`, so a new volatile header, credential field name, or generated-id shape is one edit there. Record starts fresh every time: it wipes the previous bundle (refusing to wipe a directory that is not a bundle) and never reads it. A replay bundle whose manifest is older than seven days hard-fails at collection time naming the bundle's age, so replay can never certify against fixtures that have drifted more than a week from the live proxy +A bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`) is a directory: `manifest.json` carries the record timestamp, harness git version, and format version, and each test gets a subdirectory holding one JSON file per provider call in call order (`0000-post-openai-v1-chat-completions.json`). Request headers are never stored (provider credentials never touch disk), non-JSON request bodies store a canonicalized sha256 digest instead of the bytes, and responses store status, filtered headers, and the verbatim body base64-encoded, which is part of why bundles are gitignored. `fixture_bundle.py` owns the format. Record serves the proxy the same filtered stored response replay will serve later, so the two modes are byte-identical from the proxy's side of the socket -Deliberately not here yet: streaming chunk fidelity (LIT-5742) and scoping record/replay to provider-bound traffic (LIT-5745) +Replay matches calls per test by canonical key: `fixture_canonical.py` canonicalizes the recorded request (volatile headers and credential fields out, unique markers, generated ids, uuids, and timestamps replaced with fixed placeholders, object keys sorted) and the key is the method, edge path, and a content hash, so identity survives re-records and machine changes while any real content drift comes back as an HTTP 599 naming the computed key, the closest recorded key with its file, and a content diff, and never falls through to a live call. Matching is order-independent across distinct keys (concurrent calls may interleave) and FIFO within one key (a retry loop replays its responses in recorded order); a passed test must also consume its whole recording, or teardown fails it naming a leftover key. Either way the fix is always to re-record with `E2E_FIXTURE_MODE=record`. Every rewrite rule lives in `fixture_canonical.py`, so a new volatile header, credential field name, or generated-id shape is one edit there. Record starts fresh every time: it wipes the previous bundle (refusing to wipe a directory that is not a bundle) and never reads it. A replay bundle whose manifest is older than seven days hard-fails at collection time naming the bundle's age, so replay can never certify against fixtures that have drifted more than a week from the live providers + +A replayed response carries the recorded provider response id, and `LiteLLM_SpendLogs.request_id` (the table's primary key) is that id, so a replay against a database that still holds the record run's rows silently dedupes its spend inserts and any spend assertion goes red with zero matching rows and nothing in the proxy log. Run both modes with `E2E_RESET_SPEND_LOGS=1` (plus `DATABASE_URL` in the runner env) so each session truncates the table after itself, or replay against a fresh database, which is the CI shape + +Current limits: streaming chunk fidelity is LIT-5742 (a streamed response records as one buffered body), CI wiring is LIT-5748, Bedrock cannot be mounted (SigV4 signs the Host header, so a rewritten api_base fails signature verification), multipart uploads have per-run random boundaries (the digest changes every run, so they always miss), and deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base) ## Typing diff --git a/tests/e2e/CONTRIBUTING.md b/tests/e2e/CONTRIBUTING.md index 67da1be9562..9096050a45a 100644 --- a/tests/e2e/CONTRIBUTING.md +++ b/tests/e2e/CONTRIBUTING.md @@ -54,14 +54,16 @@ Some suites need extra services the bare proxy does not start. The `logging/` OT ### Record and replay -`E2E_FIXTURE_MODE=record` runs a suite against the live proxy as usual while writing every request/response pair to a fixture bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`); `E2E_FIXTURE_MODE=replay` then runs the same suite entirely from that bundle, with no proxy traffic and no provider spend; the proxy liveness gate is skipped, so replay runs with no proxy up at all. Unset (or `live`) behaves exactly as before the knob existed +Record/replay scopes to the proxy's provider-bound traffic only. In `E2E_FIXTURE_MODE=record` the harness boots a local provider-edge server, edge-wired tests register their deployments with an `api_base` pointing at it, and every provider call the proxy makes is forwarded verbatim and written to a fixture bundle (default `tests/e2e/.fixtures`, override with `E2E_FIXTURE_DIR`). `E2E_FIXTURE_MODE=replay` runs the same tests against the same live proxy and database, but the edge answers the proxy's provider calls from the bundle instead of the provider, so the run makes zero provider calls and spends nothing while key auth, routing, cost calculation, and spend-log writes all still execute for real. Unset (or `live`) behaves exactly as before the knob existed. Both record and replay need the proxy up; only the provider is taken out of the loop ```bash -E2E_FIXTURE_MODE=record uv run pytest tests/e2e/llm_translation/ -v -E2E_FIXTURE_MODE=replay uv run pytest tests/e2e/llm_translation/ -v +E2E_FIXTURE_MODE=record uv run pytest tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py -v +E2E_FIXTURE_MODE=replay uv run pytest tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py -v ``` -Replay fails hard (`ReplayMiss`) when the tests drift from the recording, and a bundle older than seven days fails at collection time naming its age; either way the fix is to re-record. See `CLAUDE.md` in this directory for the bundle format and the transport seam +One sharp edge: a replayed response reuses the recorded provider response id, and that id is the primary key of `LiteLLM_SpendLogs`, so replaying against a database that still holds the record run's rows silently dedupes the spend writes and a spend assertion fails with zero rows. Run both commands above with `E2E_RESET_SPEND_LOGS=1` (and `DATABASE_URL` set in the pytest env) so each session truncates the spend log table after itself, or point replay at a fresh database + +Replay answers any provider call that drifted from the recording with an HTTP 599 whose body names the computed and closest recorded keys, so the test fails loudly instead of silently going live, and a bundle older than seven days fails at collection time naming its age; either way the fix is to re-record. Only tests that register edge-wired deployments participate: everything else hits its provider live in every mode, so record exactly the suite you replay. If the proxy runs in a container, set `E2E_PROVIDER_EDGE_ADVERTISE_HOST` (e.g. `host.docker.internal`) so the api_base the proxy stores can reach the edge on the pytest host, and `E2E_PROVIDER_EDGE_BIND_HOST=0.0.0.0` so the edge accepts it. See `CLAUDE.md` in this directory for the bundle format, the edge design, and the current limits (streaming, Bedrock, multipart) Tests marked `@pytest.mark.e2e` hard-fail when no proxy answers `/health/liveliness`, so a run that goes red with `No live proxy` at setup means the proxy isn't up; they never skip for a missing proxy, so an absent proxy can't be mistaken for a pass diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index da2a7da0bfa..dbe2d6e514e 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -23,12 +23,8 @@ import requests from e2e_config import CONTROL_PLANE_BASE_URL, FIXTURE_DIR, FIXTURE_MODE_RAW, PROXY_BASE_URL from e2e_db import RESET_OPT_IN_ENV, reset_spend_logs, run_spend_log_cleanup -from fixture_transport import ( - fixture_mode_collection_error, - fixture_report_lines, - parse_fixture_mode, - replay_leftover_error, -) +from fixture_mode import fixture_mode_collection_error, fixture_report_lines +from provider_edge import replay_leftover_error from junit_properties import attach_result_properties from lifecycle import ProxyClientProvider, ResourceManager from proxy_client import ProxyClient, build_proxy_client @@ -114,12 +110,10 @@ def _proxy_fail_reason() -> str | None: def pytest_runtest_setup(item: pytest.Item) -> None: """Hard-fail `e2e`-marked tests unless a proxy answers its liveness probe. Unmarked tests (unit coverage of the harness) don't touch the proxy, so they - run even when none is up. Never skip for a missing proxy. Replay mode serves - every call from the fixture bundle, so it needs no live proxy either.""" + run even when none is up. Never skip for a missing proxy. Replay mode needs + the proxy too: only provider-bound traffic replays from the bundle.""" if item.get_closest_marker("e2e") is None: return - if parse_fixture_mode(FIXTURE_MODE_RAW) == "replay": - return reason = _proxy_fail_reason() if reason is not None: pytest.fail(reason) diff --git a/tests/e2e/coverage_registry/quota_management.yaml b/tests/e2e/coverage_registry/quota_management.yaml index 4b8aa1da002..42a075681e0 100644 --- a/tests/e2e/coverage_registry/quota_management.yaml +++ b/tests/e2e/coverage_registry/quota_management.yaml @@ -47,3 +47,10 @@ - {id: quota_management.spend_tracking.failure.writes_failure_row, module: quota_management, tier: P1, behavior: spend_tracking, variant: failure, assertions: [writes_failure_row], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_log_error_logger.py", rationale: "A failed call writes a failure-status spend row"} - {id: quota_management.spend_tracking.spend_calculate.returns_cost, module: quota_management, tier: P2, behavior: spend_tracking, variant: spend_calculate, assertions: [returns_cost], exercised_on: [spend_calculate], source: "proxy/spend_tracking/spend_management_endpoints.py", rationale: "/spend/calculate prices a hypothetical request at nonzero cost"} - {id: quota_management.spend_tracking.pagination.keeps_total, module: quota_management, tier: P2, behavior: spend_tracking, variant: pagination, assertions: [keeps_total], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_management_endpoints.py", rationale: "Spend-logs v2 pagination caps page size without losing the total"} +- {id: quota_management.spend_tracking.cache_write.bills_cache_creation_rate, module: quota_management, tier: P1, behavior: spend_tracking, variant: cache_write, assertions: [bills_cache_creation_rate], exercised_on: [chat_completions], source: "litellm_core_utils/llm_cost_calc/utils.py", rationale: "OpenAI cache-write tokens land on the spend row as cache-creation tokens billed at the cache-creation rate, not silently at the input rate (#34046)"} +- {id: quota_management.spend_tracking.cost_breakdown.reports_component_costs, module: quota_management, tier: P1, behavior: spend_tracking, variant: cost_breakdown, assertions: [reports_component_costs], exercised_on: [chat_completions], source: "proxy/spend_tracking/spend_tracking_utils.py", rationale: "The spend row's metadata.cost_breakdown itemizes cache-read, cache-creation, output, and reasoning costs at the deployment's own rates and they sum to the row's spend (#31686)"} +- {id: quota_management.spend_tracking.stream_cache_read.bills_cache_read_rate, module: quota_management, tier: P1, behavior: spend_tracking, variant: stream_cache_read, assertions: [bills_cache_read_rate], exercised_on: [chat_completions], source: "litellm_core_utils/streaming_chunk_builder_utils.py", rationale: "A streamed call's reassembled usage keeps the cached-token detail so cache reads bill at the cache-read discount, not full input price (#34812)"} +- {id: quota_management.spend_tracking.messages_bridge.keeps_cache_tokens, module: quota_management, tier: P1, behavior: spend_tracking, variant: messages_bridge, assertions: [keeps_cache_tokens], exercised_on: [messages], source: "llms/anthropic/experimental_pass_through/responses_adapters/handler.py", rationale: "A /v1/messages request served by a Responses-only OpenAI model keeps its cache-read tokens and their discounted billing across the bridge (#34957)"} +- {id: quota_management.spend_tracking.service_tier.bills_tier_rates, module: quota_management, tier: P1, behavior: spend_tracking, variant: service_tier, assertions: [bills_tier_rates], exercised_on: [chat_completions], source: "cost_calculator.py", rationale: "A priority service_tier call bills input, output, and reasoning at the deployment's *_priority rates and records the tier on the row (#35923, #35925)"} +- {id: quota_management.spend_tracking.cost_headers.additive_components, module: quota_management, tier: P1, behavior: spend_tracking, variant: cost_headers, assertions: [additive_components], exercised_on: [chat_completions], source: "proxy/common_request_processing.py", rationale: "The x-litellm-response-cost-* component headers sum to the total, input covers only fresh tokens, and reasoning stays a subset of output (#36965)"} +- {id: quota_management.spend_tracking.passthrough_stream.injects_usage_cost, module: quota_management, tier: P1, behavior: spend_tracking, variant: passthrough_stream, assertions: [injects_usage_cost], exercised_on: [openai_passthrough], source: "proxy/pass_through_endpoints/streaming_handler.py", rationale: "With include_cost_in_streaming_usage on, the /openai passthrough's final streaming usage frame carries the proxy-computed cost (#36503). Uncovered: the flag is only settable in litellm_settings, and the shared e2e stack does not turn it on yet"} diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index a5c3729f4be..8bf39f6021f 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -13,7 +13,8 @@ from pathlib import Path from dotenv import load_dotenv -from fixture_transport import deterministic_marker, parse_fixture_mode +from fixture_mode import deterministic_marker, parse_fixture_mode +from provider_edge import provider_edge_api_base # Local runs keep provider / DataDog keys in tests/e2e/.env (see CONTRIBUTING.md). # Compose injects them into the proxy container, but pytest on the host does not @@ -92,15 +93,24 @@ PROPAGATION_TIMEOUT = float(os.environ.get("E2E_PROPAGATION_TIMEOUT", "15")) EXPECT_RUST = os.environ.get("E2E_EXPECT_RUST", "").strip().lower() in ("1", "true", "yes") -# Record/replay fixture selection (see fixture_transport.py). The raw mode value -# is parsed and validated there; "live" (the default, also for empty values) -# means the harness behaves exactly as before this knob existed. +# Record/replay fixture selection (see fixture_mode.py and provider_edge.py). +# The raw mode value is parsed and validated there; "live" (the default, also +# for empty values) means the harness behaves exactly as before this knob +# existed. FIXTURE_MODE_RAW = os.environ.get("E2E_FIXTURE_MODE", "live") FIXTURE_DIR = Path( os.environ.get("E2E_FIXTURE_DIR", "").strip() or str(Path(__file__).resolve().parent / ".fixtures") ) +# Where the provider-edge server binds, and the host name edge api_base URLs +# advertise to the proxy. They differ when the proxy runs in a container and +# reaches the pytest host via a gateway name like host.docker.internal. +PROVIDER_EDGE_BIND_HOST = os.environ.get("E2E_PROVIDER_EDGE_BIND_HOST", "").strip() or "127.0.0.1" +PROVIDER_EDGE_ADVERTISE_HOST = ( + os.environ.get("E2E_PROVIDER_EDGE_ADVERTISE_HOST", "").strip() or PROVIDER_EDGE_BIND_HOST +) + # Deliberately modest concurrency. The suite shares its proxy with every other # suite in the run, and 750 users at spawn rate 50 saturated the request path hard # enough to distort latency-sensitive neighbours (and to spend real provider money @@ -157,6 +167,20 @@ def datadog_mcp_url(*, toolsets: str = "core") -> str: return f"{base}?toolsets={toolsets}" if toolsets else base +def provider_edge_base(mount: str) -> str | None: + """The api_base an edge-wired deployment should register with, using this + process's fixture-mode and edge-host configuration: None in live mode, the + shared edge server's mount URL in record and replay.""" + return provider_edge_api_base( + mount, + mode_raw=FIXTURE_MODE_RAW, + bundle_dir=FIXTURE_DIR, + bind_host=PROVIDER_EDGE_BIND_HOST, + advertise_host=PROVIDER_EDGE_ADVERTISE_HOST, + forward_timeout=REQUEST_TIMEOUT, + ) + + def unique_marker() -> str: """A short unique token per call/run, so concurrent runs and the shared response cache never collide on prompts, tags, or customer ids. In record diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index cb6fc7a01e5..03f201e946e 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -647,3 +647,37 @@ def download( content_type=_hdr(resp, "content-type"), body=resp.text, ) + + +class RawResponse(BaseModel): + """A verbatim upstream HTTP response for the provider edge (provider_edge.py): + status, lowercased headers, raw bytes. No Result classification because the + edge relays provider errors to the proxy untouched.""" + + status_code: int + headers: dict[str, str] + body: bytes + + +def forward( + method: str, + url: str, + *, + headers: dict[str, str], + body: bytes | None, + timeout: float = 60.0, +) -> RawResponse | NetworkError: + """Relay one provider-bound request verbatim for the provider edge's record + mode. No retries, no redirects, no schema: the proxy owns retry policy and + the recorded bundle must hold exactly what the provider returned.""" + try: + resp = requests.request( + method, url, headers=headers, data=body, timeout=timeout, allow_redirects=False + ) + except requests.RequestException as exc: + return NetworkError(message=str(exc)) + return RawResponse( + status_code=resp.status_code, + headers={name.lower(): value for name, value in resp.headers.items()}, + body=resp.content, + ) diff --git a/tests/e2e/fixture_bundle.py b/tests/e2e/fixture_bundle.py index 615ae8df1a4..6feb40fc8bc 100644 --- a/tests/e2e/fixture_bundle.py +++ b/tests/e2e/fixture_bundle.py @@ -1,17 +1,18 @@ -"""On-disk fixture bundle format for record/replay e2e runs (LIT-5729). +"""On-disk fixture bundle format for record/replay e2e runs (LIT-5729/LIT-5745). A bundle is a directory: one ``manifest.json`` (record timestamp + harness version + format version) plus one subdirectory per test, holding one JSON file -per transport interaction in call order. Bundles older than +per provider-bound interaction in call order. Bundles older than ``MAX_BUNDLE_AGE`` hard-fail replay at collection time (see conftest), so a green replay run can never certify against fixtures that have drifted more than -a week from the live proxy. +a week from the live providers. -This module owns the format only. The transports that produce and consume it -live in fixture_transport.py and the canonical match keys they compute live in -fixture_canonical.py (LIT-5741); streaming chunk fidelity and provider-scoping -are follow-ups (LIT-5742/5745). Every interaction file stores the full redacted -request because replay matches on its canonicalized content. +This module owns the format only. The provider-edge server that produces and +consumes it lives in provider_edge.py (LIT-5745) and the canonical match keys +it computes live in fixture_canonical.py (LIT-5741); streaming chunk fidelity +is a follow-up (LIT-5742). Every interaction file stores the full redacted +request because replay matches on its canonicalized content, and the response +as the raw HTTP status, filtered headers, and base64 body the provider sent. """ from __future__ import annotations @@ -23,29 +24,14 @@ import subprocess from dataclasses import dataclass, field from datetime import datetime, timedelta, timezone from pathlib import Path -from typing import Annotated, Final, Literal +from typing import Final -from pydantic import BaseModel, Field, JsonValue, TypeAdapter +from pydantic import BaseModel, JsonValue -from e2e_http import ( - BinaryStream, - NetworkError, - ProbeResult, - RateLimitedError, - Result, - StreamingResponse, - Success, - UnauthorizedError, - UnknownApiError, - ValidationError, -) - -BUNDLE_FORMAT_VERSION: Final = 1 +BUNDLE_FORMAT_VERSION: Final = 2 MAX_BUNDLE_AGE: Final = timedelta(days=7) MANIFEST_FILENAME: Final = "manifest.json" -_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) - class Manifest(BaseModel): format_version: int @@ -54,13 +40,14 @@ class Manifest(BaseModel): class RecordedRequest(BaseModel): - """The request as the transport saw it, auth header values and credential - body/form fields redacted. + """The provider-bound request as the edge saw it, headers empty (SDK + telemetry headers vary run to run and auth material never touches disk). Replay matches on the canonical content key fixture_canonical.py computes - over ``method`` (the transport verb, not the HTTP verb), ``path``, and the - canonicalized headers, params, body, form, and file identity. File uploads - store a content digest instead of the bytes.""" + over ``method``, ``path`` (the edge path including the provider mount, + query string excluded), and the canonicalized headers, params, body, form, + and file identity. Non-JSON bodies store a canonicalized content digest + instead of the bytes.""" method: str path: str @@ -73,85 +60,19 @@ class RecordedRequest(BaseModel): file_bytes: int | None = None -class RecordedResult(BaseModel): - """A ``Result[R]`` flattened for disk. ``data`` holds the success payload as - raw JSON; replay re-validates it against the ``response_type`` the caller - passes, exactly like a live response body.""" +class RecordedHttpResponse(BaseModel): + """The provider's raw HTTP response: status, headers minus hop-by-hop and + volatile entries (see provider_edge.py), and the body as base64 so binary + payloads survive JSON.""" - shape: Literal["result"] = "result" - kind: Literal["success", "network", "unauthorized", "rate_limited", "validation", "unknown"] - status_code: int | None = None - data: JsonValue | None = None - message: str | None = None - body: str | None = None - retry_after_seconds: int | None = None - - -class RecordedStreaming(BaseModel): - shape: Literal["streaming"] = "streaming" - payload: StreamingResponse - - -class RecordedBinary(BaseModel): - shape: Literal["binary"] = "binary" - payload: BinaryStream - - -class RecordedProbe(BaseModel): - shape: Literal["probe"] = "probe" - payload: ProbeResult - - -type RecordedResponse = RecordedResult | RecordedStreaming | RecordedBinary | RecordedProbe + status_code: int + headers: dict[str, str] + body_b64: str class Interaction(BaseModel): request: RecordedRequest - response: Annotated[ - RecordedResult | RecordedStreaming | RecordedBinary | RecordedProbe, - Field(discriminator="shape"), - ] - - -def to_json_value(model: BaseModel) -> JsonValue: - return _JSON.validate_json(model.model_dump_json(by_alias=True)) - - -def from_result[R: BaseModel](result: Result[R]) -> RecordedResult: - match result: - case Success(status_code=status_code, data=data): - return RecordedResult(kind="success", status_code=status_code, data=to_json_value(data)) - case NetworkError(message=message): - return RecordedResult(kind="network", message=message) - case UnauthorizedError(): - return RecordedResult(kind="unauthorized") - case RateLimitedError(retry_after_seconds=retry_after_seconds, body=body): - return RecordedResult(kind="rate_limited", retry_after_seconds=retry_after_seconds, body=body) - case ValidationError(message=message): - return RecordedResult(kind="validation", message=message) - case UnknownApiError(status_code=status_code, body=body): - return RecordedResult(kind="unknown", status_code=status_code, body=body) - - -def to_result[R: BaseModel](recorded: RecordedResult, response_type: type[R]) -> Result[R]: - match recorded.kind: - case "success": - return Success( - status_code=recorded.status_code or 200, - data=response_type.model_validate(recorded.data), - ) - case "network": - return NetworkError(message=recorded.message or "") - case "unauthorized": - return UnauthorizedError() - case "rate_limited": - return RateLimitedError( - retry_after_seconds=recorded.retry_after_seconds, body=recorded.body or "" - ) - case "validation": - return ValidationError(message=recorded.message or "") - case "unknown": - return UnknownApiError(status_code=recorded.status_code or 0, body=recorded.body or "") + response: RecordedHttpResponse def slugify(raw: str, *, limit: int = 60) -> str: @@ -198,7 +119,7 @@ class BundleRecorder: root: Path _ordinals: dict[str, int] = field(default_factory=dict) - def record(self, *, test_key: str, request: RecordedRequest, response: RecordedResponse) -> None: + def record(self, *, test_key: str, request: RecordedRequest, response: RecordedHttpResponse) -> None: slug = slug_for_test(test_key) ordinal = self._ordinals.get(slug, 0) self._ordinals[slug] = ordinal + 1 diff --git a/tests/e2e/fixture_mode.py b/tests/e2e/fixture_mode.py new file mode 100644 index 00000000000..110f44380b4 --- /dev/null +++ b/tests/e2e/fixture_mode.py @@ -0,0 +1,132 @@ +"""Fixture-mode selection and per-test determinism for record/replay e2e runs. + +``E2E_FIXTURE_MODE`` is live (the default; nothing changes), record, or replay. +This module owns everything mode-shaped that is independent of the provider +edge itself: parsing the raw env value, the collection-time gate that aborts a +run whose mode can never work (unknown value, or replay against a missing or +stale bundle), the pytest report-header lines, the running test's node id, and +the deterministic per-test marker that lets a replay run regenerate exactly +the requests the record run sent. The provider-edge server that records and +serves provider traffic lives in provider_edge.py (LIT-5745). +""" + +from __future__ import annotations + +import hashlib +import os +from dataclasses import dataclass +from datetime import datetime +from pathlib import Path +from typing import Final, Literal, assert_never + +from fixture_bundle import ( + FreshBundle, + StaleBundle, + UnreadableBundle, + check_freshness, + format_age, +) + +type FixtureMode = Literal["live", "record", "replay"] + +FIXTURE_MODES: Final[tuple[FixtureMode, ...]] = ("live", "record", "replay") + +SESSION_TEST_KEY: Final = "session" + + +@dataclass(frozen=True, slots=True) +class InvalidFixtureMode: + value: str + + +def parse_fixture_mode(raw: str) -> FixtureMode | InvalidFixtureMode: + normalized = raw.strip().lower() or "live" + match normalized: + case "live" | "record" | "replay": + return normalized + case _: + return InvalidFixtureMode(value=raw) + + +def current_test_key() -> str: + """The pytest node id of the running test, from the PYTEST_CURRENT_TEST env + var pytest maintains (`` (setup|call|teardown)``); ``session`` for + calls outside any test (e.g. session-finish cleanup).""" + raw = os.environ.get("PYTEST_CURRENT_TEST", "") + if not raw: + return SESSION_TEST_KEY + return raw.rsplit(" (", 1)[0] + + +class ReplayMiss(AssertionError): + """Replay had no recorded interaction for a provider call the proxy made. + The suite drifted from the bundle (or the bundle from the suite): re-record.""" + + +_marker_ordinals: Final[dict[str, int]] = {} + + +def deterministic_marker() -> str: + """Stable stand-in for uuid-based unique markers in record and replay modes: + the Nth marker of a test is a pure function of the test's node id and N, so a + replay run regenerates exactly the model names, prompts, and tags the record + run sent and every recorded provider interaction still matches its key.""" + test_key = current_test_key() + ordinal = _marker_ordinals.get(test_key, 0) + _marker_ordinals[test_key] = ordinal + 1 + return hashlib.sha1(f"{test_key}#{ordinal}".encode()).hexdigest()[:12] + + +def fixture_mode_collection_error(mode_raw: str, bundle_dir: Path, *, now: datetime) -> str | None: + """Session-abort reason for a fixture-mode setup that can never work, or None. + Called at collection time (conftest pytest_sessionstart) so a stale or missing + bundle fails the whole run up front, naming the bundle age, instead of failing + every test individually.""" + mode = parse_fixture_mode(mode_raw) + match mode: + case InvalidFixtureMode(value=value): + return f"E2E_FIXTURE_MODE={value!r} is not one of {', '.join(FIXTURE_MODES)}" + case "live" | "record": + return None + case "replay": + freshness = check_freshness(bundle_dir, now=now) + match freshness: + case FreshBundle(): + return None + case StaleBundle(recorded_at=recorded_at, age=age, limit=limit): + return ( + f"fixture bundle at {bundle_dir} is stale: recorded {recorded_at.isoformat()}, " + f"age {format_age(age)} exceeds the {limit.days}-day limit; " + "re-record with E2E_FIXTURE_MODE=record" + ) + case UnreadableBundle(reason=reason): + return f"E2E_FIXTURE_MODE=replay cannot use bundle at {bundle_dir}: {reason}" + case _: + assert_never(freshness) + case _: + assert_never(mode) + + +def fixture_report_lines(mode_raw: str, bundle_dir: Path, *, now: datetime) -> list[str]: + """pytest report-header lines; empty in live mode so an unset + E2E_FIXTURE_MODE keeps today's output byte-identical.""" + mode = parse_fixture_mode(mode_raw) + match mode: + case InvalidFixtureMode() | "live": + return [] + case "record": + return [f"e2e fixture mode: record -> {bundle_dir}"] + case "replay": + freshness = check_freshness(bundle_dir, now=now) + match freshness: + case FreshBundle(manifest=manifest): + return [ + f"e2e fixture mode: replay <- {bundle_dir} " + f"(recorded {manifest.recorded_at.isoformat()}, harness {manifest.harness_version})" + ] + case StaleBundle() | UnreadableBundle(): + return [f"e2e fixture mode: replay <- {bundle_dir}"] + case _: + assert_never(freshness) + case _: + assert_never(mode) diff --git a/tests/e2e/fixture_transport.py b/tests/e2e/fixture_transport.py deleted file mode 100644 index ce4eec701ca..00000000000 --- a/tests/e2e/fixture_transport.py +++ /dev/null @@ -1,724 +0,0 @@ -"""Record/replay transports behind the same ``Transport`` protocol (LIT-5729). - -``RecordingTransport`` decorates the live transport: every call passes through -unchanged and its request/response pair is appended to the fixture bundle. -``ReplayTransport`` implements the protocol from a recorded bundle alone: no -HTTP, no proxy, no provider spend. Because both fulfil ``Transport``, no test -or client changes shape; ``build_proxy_client`` picks the transport from -``E2E_FIXTURE_MODE`` (live | record | replay, default live). - -Replay matches each call by test node id and canonical content key -(fixture_canonical.py, LIT-5741): volatile headers, credential fields, unique -markers, generated ids, and timestamps are canonicalized out before hashing, so -matching is order-independent across distinct keys, FIFO within a key, and a -miss fails hard (``ReplayMiss``) printing the computed key and the closest -recorded key without ever falling through to a live call. Streaming chunk -fidelity is LIT-5742; scoping record/replay to provider-bound traffic is -LIT-5745. -""" - -from __future__ import annotations - -import difflib -import functools -import hashlib -import os -from collections import deque -from dataclasses import dataclass, field -from datetime import datetime -from itertools import islice -from pathlib import Path -from typing import Final, Literal, assert_never - -from pydantic import BaseModel, JsonValue - -from e2e_http import AuthHeaders, BinaryStream, ProbeResult, Result, StreamingResponse -from fixture_bundle import ( - BundleRecorder, - FreshBundle, - Interaction, - LoadedBundle, - RecordedBinary, - RecordedProbe, - RecordedRequest, - RecordedResponse, - RecordedResult, - RecordedStreaming, - StaleBundle, - UnreadableBundle, - UnsafeBundleDir, - check_freshness, - format_age, - from_result, - interaction_filename, - load_bundle, - prepare_bundle, - slug_for_test, - to_json_value, - to_result, -) -from fixture_canonical import CanonicalRequest, canonicalize, is_secret_field -from transport import Transport - -type FixtureMode = Literal["live", "record", "replay"] - -FIXTURE_MODES: Final[tuple[FixtureMode, ...]] = ("live", "record", "replay") - -SESSION_TEST_KEY: Final = "session" - -REDACTED_HEADER_NAMES: Final[frozenset[str]] = frozenset({"authorization", "x-litellm-api-key"}) -REDACTED_VALUE: Final = "" - - -@dataclass(frozen=True, slots=True) -class InvalidFixtureMode: - value: str - - -def parse_fixture_mode(raw: str) -> FixtureMode | InvalidFixtureMode: - normalized = raw.strip().lower() or "live" - match normalized: - case "live" | "record" | "replay": - return normalized - case _: - return InvalidFixtureMode(value=raw) - - -def current_test_key() -> str: - """The pytest node id of the running test, from the PYTEST_CURRENT_TEST env - var pytest maintains (`` (setup|call|teardown)``); ``session`` for - calls outside any test (e.g. session-finish cleanup).""" - raw = os.environ.get("PYTEST_CURRENT_TEST", "") - if not raw: - return SESSION_TEST_KEY - return raw.rsplit(" (", 1)[0] - - -class ReplayMiss(AssertionError): - """Replay had no recorded interaction for a call the suite made. The test - drifted from the bundle (or the bundle from the suite): re-record.""" - - -_marker_ordinals: Final[dict[str, int]] = {} - - -def deterministic_marker() -> str: - """Stable stand-in for uuid-based unique markers in record and replay modes: - the Nth marker of a test is a pure function of the test's node id and N, so a - replay run regenerates exactly the model names, prompts, and tags the record - run sent and every recorded poll response still satisfies its predicate.""" - test_key = current_test_key() - ordinal = _marker_ordinals.get(test_key, 0) - _marker_ordinals[test_key] = ordinal + 1 - return hashlib.sha1(f"{test_key}#{ordinal}".encode()).hexdigest()[:12] - - -def _dump_flat(model: BaseModel | None) -> dict[str, str]: - if model is None: - return {} - dumped: dict[str, object] = model.model_dump(by_alias=True, exclude_none=True) - return {key: str(value) for key, value in dumped.items()} - - -def _redact(headers: dict[str, str]) -> dict[str, str]: - return { - name: REDACTED_VALUE if name.lower() in REDACTED_HEADER_NAMES else value - for name, value in headers.items() - } - - -def _redact_secret_fields(value: JsonValue) -> JsonValue: - match value: - case dict(): - return { - key: REDACTED_VALUE - if is_secret_field(key) and item is not None - else _redact_secret_fields(item) - for key, item in value.items() - } - case list(): - return [_redact_secret_fields(item) for item in value] - case _: - return value - - -def _redact_flat(fields: dict[str, str]) -> dict[str, str]: - return { - key: REDACTED_VALUE if is_secret_field(key) else value for key, value in fields.items() - } - - -def recorded_request( - method: str, - path: str, - *, - headers: BaseModel, - body: BaseModel | None = None, - params: BaseModel | None = None, - form: BaseModel | None = None, - file_name: str | None = None, - file_content: bytes | None = None, -) -> RecordedRequest: - return RecordedRequest( - method=method, - path=path, - headers=_redact(_dump_flat(headers)), - params=_redact_flat(_dump_flat(params)), - body=None if body is None else _redact_secret_fields(to_json_value(body)), - form=None if form is None else _redact_flat(_dump_flat(form)), - file_name=file_name, - file_sha256=None if file_content is None else hashlib.sha256(file_content).hexdigest(), - file_bytes=None if file_content is None else len(file_content), - ) - - -@dataclass(frozen=True, slots=True) -class RecordingTransport: - """Decorator over the live transport: forwards every call and appends the - interaction to the bundle, so a green live run leaves behind exactly the - traffic replay needs.""" - - inner: Transport - recorder: BundleRecorder - - def _record(self, request: RecordedRequest, response: RecordedResponse) -> None: - self.recorder.record(test_key=current_test_key(), request=request, response=response) - - def bearer(self, key: str) -> AuthHeaders: - return self.inner.bearer(key) - - @property - def master(self) -> AuthHeaders: - return self.inner.master - - def post[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - result = self.inner.post(path, headers=headers, json=json, response_type=response_type) - self._record(recorded_request("post", path, headers=headers, body=json), from_result(result)) - return result - - def get[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - params: BaseModel, - response_type: type[R], - timeout: float | None = None, - ) -> Result[R]: - result = self.inner.get( - path, headers=headers, params=params, response_type=response_type, timeout=timeout - ) - self._record(recorded_request("get", path, headers=headers, params=params), from_result(result)) - return result - - def delete[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - json: BaseModel, - response_type: type[R], - params: BaseModel | None = None, - ) -> Result[R]: - result = self.inner.delete( - path, headers=headers, json=json, response_type=response_type, params=params - ) - self._record( - recorded_request("delete", path, headers=headers, body=json, params=params), - from_result(result), - ) - return result - - def patch[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - result = self.inner.patch(path, headers=headers, json=json, response_type=response_type) - self._record(recorded_request("patch", path, headers=headers, body=json), from_result(result)) - return result - - def put[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - result = self.inner.put(path, headers=headers, json=json, response_type=response_type) - self._record(recorded_request("put", path, headers=headers, body=json), from_result(result)) - return result - - def stream(self, path: str, *, headers: BaseModel, json: BaseModel) -> StreamingResponse: - response = self.inner.stream(path, headers=headers, json=json) - self._record( - recorded_request("stream", path, headers=headers, body=json), - RecordedStreaming(payload=response), - ) - return response - - def stream_binary( - self, path: str, *, headers: BaseModel, json: BaseModel, chunk_size: int = 8192 - ) -> BinaryStream: - response = self.inner.stream_binary(path, headers=headers, json=json, chunk_size=chunk_size) - self._record( - recorded_request("stream_binary", path, headers=headers, body=json), - RecordedBinary(payload=response), - ) - return response - - def send( - self, - path: str, - *, - headers: BaseModel, - json: BaseModel, - params: BaseModel | None = None, - stream: bool = False, - ) -> StreamingResponse: - response = self.inner.send(path, headers=headers, json=json, params=params, stream=stream) - self._record( - recorded_request("send", path, headers=headers, body=json, params=params), - RecordedStreaming(payload=response), - ) - return response - - def probe(self, path: str, *, params: BaseModel) -> ProbeResult: - response = self.inner.probe(path, params=params) - self._record( - recorded_request("probe", path, headers=self.master, params=params), - RecordedProbe(payload=response), - ) - return response - - def upload[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - form: BaseModel, - filename: str, - content: bytes, - file_content_type: str = "application/jsonl", - file_field: str = "file", - params: BaseModel | None = None, - response_type: type[R], - ) -> Result[R]: - result = self.inner.upload( - path, - headers=headers, - form=form, - filename=filename, - content=content, - file_content_type=file_content_type, - file_field=file_field, - params=params, - response_type=response_type, - ) - self._record( - recorded_request( - "upload", - path, - headers=headers, - params=params, - form=form, - file_name=filename, - file_content=content, - ), - from_result(result), - ) - return result - - def download(self, path: str, *, headers: BaseModel) -> StreamingResponse: - response = self.inner.download(path, headers=headers) - self._record( - recorded_request("download", path, headers=headers), - RecordedStreaming(payload=response), - ) - return response - - -def _build_pool(recorded: tuple[Interaction, ...]) -> dict[str, deque[Interaction]]: - keys: Final = tuple(canonicalize(interaction.request).key for interaction in recorded) - return { - key: deque( - interaction - for candidate_key, interaction in zip(keys, recorded, strict=True) - if candidate_key == key - ) - for key in dict.fromkeys(keys) - } - - -def _closest_recorded( - canonical: CanonicalRequest, recorded: tuple[Interaction, ...] -) -> tuple[CanonicalRequest, str]: - candidates: Final = tuple(canonicalize(interaction.request) for interaction in recorded) - ratios: Final = tuple( - difflib.SequenceMatcher( - None, f"{canonical.method} {canonical.path}\n{canonical.content}", - f"{candidate.method} {candidate.path}\n{candidate.content}", - ).ratio() - for candidate in candidates - ) - best: Final = max(range(len(candidates)), key=lambda index: ratios[index]) - return candidates[best], interaction_filename(best, recorded[best].request) - - -def _miss_message(test_key: str, slug: str, canonical: CanonicalRequest, bundle: LoadedBundle) -> str: - recorded: Final = bundle.interactions.get(slug, ()) - if not recorded: - return ( - f"replay miss for {test_key}: computed key {canonical.key} but nothing is recorded " - f"under {slug}; re-record with E2E_FIXTURE_MODE=record" - ) - closest, closest_file = _closest_recorded(canonical, recorded) - diff: Final = "\n".join( - islice( - difflib.unified_diff( - closest.pretty_content().splitlines(), - canonical.pretty_content().splitlines(), - fromfile=f"closest recorded ({closest_file})", - tofile="test made", - lineterm="", - ), - 60, - ) - ) - return ( - f"replay miss for {test_key}: no recorded interaction matches key {canonical.key}; " - f"closest recorded key is {closest.key} ({closest_file})\n{diff}\n" - "re-record with E2E_FIXTURE_MODE=record" - ) - - -@dataclass(slots=True) -class ReplaySource: - """One shared pool per test over a loaded bundle, so every client built in - the session consumes the same recorded interactions. Every pool is built - once at construction and per-key consumption is a single atomic deque pop, - so concurrent replay calls never race. Calls match by canonical content - key: order-independent across distinct keys (concurrent tests interleave - calls nondeterministically), FIFO within one key (a poll loop replays its - recorded responses in recorded order).""" - - bundle: LoadedBundle - _pools: dict[str, dict[str, deque[Interaction]]] = field(init=False) - - def __post_init__(self) -> None: - self._pools = { - slug: _build_pool(recorded) for slug, recorded in self.bundle.interactions.items() - } - - def _pool(self, slug: str) -> dict[str, deque[Interaction]]: - return self._pools.get(slug, {}) - - def next_interaction(self, request: RecordedRequest) -> Interaction: - test_key: Final = current_test_key() - slug: Final = slug_for_test(test_key) - pool: Final = self._pool(slug) - canonical: Final = canonicalize(request) - queue: Final = pool.get(canonical.key) - if queue is None: - raise ReplayMiss(_miss_message(test_key, slug, canonical, self.bundle)) - try: - return queue.popleft() - except IndexError: - raise ReplayMiss( - f"replay exhausted for {test_key}: every recorded interaction for key " - f"{canonical.key} is already consumed; re-record with E2E_FIXTURE_MODE=record" - ) from None - - def leftover_error(self, test_key: str) -> str | None: - """Non-None when the test consumed fewer interactions than were recorded, - meaning a passing replay proved less than the bundle claims.""" - slug: Final = slug_for_test(test_key) - recorded: Final = self.bundle.interactions.get(slug, ()) - if not recorded: - return None - leftover: Final = tuple( - interaction for queue in self._pool(slug).values() for interaction in queue - ) - if not leftover: - return None - return ( - f"replay incomplete for {test_key}: {len(leftover)} of {len(recorded)} recorded " - f"interactions never consumed, e.g. {canonicalize(leftover[0].request).key}; " - "re-record with E2E_FIXTURE_MODE=record" - ) - - -def _expect_result(interaction: Interaction) -> RecordedResult: - match interaction.response: - case RecordedResult() as recorded: - return recorded - case RecordedStreaming() | RecordedBinary() | RecordedProbe(): - raise ReplayMiss( - f"recorded {interaction.request.method} {interaction.request.path} is not a typed result" - ) - - -def _expect_streaming(interaction: Interaction) -> StreamingResponse: - match interaction.response: - case RecordedStreaming(payload=payload): - return payload - case RecordedResult() | RecordedBinary() | RecordedProbe(): - raise ReplayMiss( - f"recorded {interaction.request.method} {interaction.request.path} is not a streaming response" - ) - - -@dataclass(frozen=True, slots=True) -class ReplayTransport: - """A ``Transport`` served entirely from a recorded bundle: never opens a - connection, so a replay run cannot bill a provider.""" - - source: ReplaySource - master_key: str - - def bearer(self, key: str) -> AuthHeaders: - return AuthHeaders(authorization=f"Bearer {key}") - - @property - def master(self) -> AuthHeaders: - return self.bearer(self.master_key) - - def post[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - return to_result( - _expect_result( - self.source.next_interaction(recorded_request("post", path, headers=headers, body=json)) - ), - response_type, - ) - - def get[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - params: BaseModel, - response_type: type[R], - timeout: float | None = None, - ) -> Result[R]: - return to_result( - _expect_result( - self.source.next_interaction(recorded_request("get", path, headers=headers, params=params)) - ), - response_type, - ) - - def delete[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - json: BaseModel, - response_type: type[R], - params: BaseModel | None = None, - ) -> Result[R]: - return to_result( - _expect_result( - self.source.next_interaction( - recorded_request("delete", path, headers=headers, body=json, params=params) - ) - ), - response_type, - ) - - def patch[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - return to_result( - _expect_result( - self.source.next_interaction(recorded_request("patch", path, headers=headers, body=json)) - ), - response_type, - ) - - def put[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - return to_result( - _expect_result( - self.source.next_interaction(recorded_request("put", path, headers=headers, body=json)) - ), - response_type, - ) - - def stream(self, path: str, *, headers: BaseModel, json: BaseModel) -> StreamingResponse: - return _expect_streaming( - self.source.next_interaction(recorded_request("stream", path, headers=headers, body=json)) - ) - - def stream_binary( - self, path: str, *, headers: BaseModel, json: BaseModel, chunk_size: int = 8192 - ) -> BinaryStream: - interaction = self.source.next_interaction( - recorded_request("stream_binary", path, headers=headers, body=json) - ) - match interaction.response: - case RecordedBinary(payload=payload): - return payload - case RecordedResult() | RecordedStreaming() | RecordedProbe(): - raise ReplayMiss( - f"recorded stream_binary {interaction.request.path} is not a binary stream" - ) - - def send( - self, - path: str, - *, - headers: BaseModel, - json: BaseModel, - params: BaseModel | None = None, - stream: bool = False, - ) -> StreamingResponse: - return _expect_streaming( - self.source.next_interaction( - recorded_request("send", path, headers=headers, body=json, params=params) - ) - ) - - def probe(self, path: str, *, params: BaseModel) -> ProbeResult: - interaction = self.source.next_interaction( - recorded_request("probe", path, headers=self.master, params=params) - ) - match interaction.response: - case RecordedProbe(payload=payload): - return payload - case RecordedResult() | RecordedStreaming() | RecordedBinary(): - raise ReplayMiss(f"recorded probe {interaction.request.path} is not a probe result") - - def upload[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - form: BaseModel, - filename: str, - content: bytes, - file_content_type: str = "application/jsonl", - file_field: str = "file", - params: BaseModel | None = None, - response_type: type[R], - ) -> Result[R]: - return to_result( - _expect_result( - self.source.next_interaction( - recorded_request( - "upload", - path, - headers=headers, - params=params, - form=form, - file_name=filename, - file_content=content, - ) - ) - ), - response_type, - ) - - def download(self, path: str, *, headers: BaseModel) -> StreamingResponse: - return _expect_streaming( - self.source.next_interaction(recorded_request("download", path, headers=headers)) - ) - - -@functools.lru_cache(maxsize=8) -def _shared_recorder(root: Path) -> BundleRecorder: - prepared = prepare_bundle(root) - if isinstance(prepared, UnsafeBundleDir): - raise ValueError(f"E2E_FIXTURE_DIR {prepared.path} {prepared.reason}") - return prepared - - -@functools.lru_cache(maxsize=8) -def _shared_replay_source(root: Path) -> ReplaySource: - loaded = load_bundle(root) - if isinstance(loaded, UnreadableBundle): - raise ValueError(f"cannot replay from {root}: {loaded.reason}") - return ReplaySource(bundle=loaded) - - -def replay_leftover_error(*, mode_raw: str, bundle_dir: Path, test_key: str) -> str | None: - """Teardown-time completeness check: in replay mode a passed test with - unconsumed recorded interactions must fail instead of passing against a - recording it no longer matches. Inert in every other mode.""" - if parse_fixture_mode(mode_raw) != "replay": - return None - return _shared_replay_source(bundle_dir).leftover_error(test_key) - - -def select_transport( - live: Transport, *, mode_raw: str, bundle_dir: Path, master_key: str -) -> Transport: - """The one seam every client build goes through: wraps (record), replaces - (replay), or passes through (live) the transport per E2E_FIXTURE_MODE. The - recorder and replay cursors are process-wide singletons per bundle dir, so - every client in a session shares one bundle and one recorded sequence.""" - mode = parse_fixture_mode(mode_raw) - match mode: - case InvalidFixtureMode(value=value): - raise ValueError(f"E2E_FIXTURE_MODE={value!r} is not one of {', '.join(FIXTURE_MODES)}") - case "live": - return live - case "record": - return RecordingTransport(inner=live, recorder=_shared_recorder(bundle_dir)) - case "replay": - return ReplayTransport(source=_shared_replay_source(bundle_dir), master_key=master_key) - case _: - assert_never(mode) - - -def fixture_mode_collection_error(mode_raw: str, bundle_dir: Path, *, now: datetime) -> str | None: - """Session-abort reason for a fixture-mode setup that can never work, or None. - Called at collection time (conftest pytest_sessionstart) so a stale or missing - bundle fails the whole run up front, naming the bundle age, instead of failing - every test individually.""" - mode = parse_fixture_mode(mode_raw) - match mode: - case InvalidFixtureMode(value=value): - return f"E2E_FIXTURE_MODE={value!r} is not one of {', '.join(FIXTURE_MODES)}" - case "live" | "record": - return None - case "replay": - freshness = check_freshness(bundle_dir, now=now) - match freshness: - case FreshBundle(): - return None - case StaleBundle(recorded_at=recorded_at, age=age, limit=limit): - return ( - f"fixture bundle at {bundle_dir} is stale: recorded {recorded_at.isoformat()}, " - f"age {format_age(age)} exceeds the {limit.days}-day limit; " - "re-record with E2E_FIXTURE_MODE=record" - ) - case UnreadableBundle(reason=reason): - return f"E2E_FIXTURE_MODE=replay cannot use bundle at {bundle_dir}: {reason}" - case _: - assert_never(freshness) - case _: - assert_never(mode) - - -def fixture_report_lines(mode_raw: str, bundle_dir: Path, *, now: datetime) -> list[str]: - """pytest report-header lines; empty in live mode so an unset - E2E_FIXTURE_MODE keeps today's output byte-identical.""" - mode = parse_fixture_mode(mode_raw) - match mode: - case InvalidFixtureMode() | "live": - return [] - case "record": - return [f"e2e fixture mode: record -> {bundle_dir}"] - case "replay": - freshness = check_freshness(bundle_dir, now=now) - match freshness: - case FreshBundle(manifest=manifest): - return [ - f"e2e fixture mode: replay <- {bundle_dir} " - f"(recorded {manifest.recorded_at.isoformat()}, harness {manifest.harness_version})" - ] - case StaleBundle() | UnreadableBundle(): - return [f"e2e fixture mode: replay <- {bundle_dir}"] - case _: - assert_never(freshness) - case _: - assert_never(mode) diff --git a/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py b/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py index 92b33fef85f..735f1a4a703 100644 --- a/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py +++ b/tests/e2e/llm_translation/test_audio_transcriptions_e2e.py @@ -3,7 +3,10 @@ Registers an OpenAI speech-to-text deployment at runtime and uploads a spoken weather question (the realtime suite's 24kHz WAV fixture) as multipart, asserting the returned transcript is non-empty and mentions the word it was asked about. -Also pins missing file/model negatives. +Also pins missing file/model negatives. A model-less request comes back as one of +two 400s depending on whether any wildcard deployment happens to be registered on +the shared proxy, so the assertion accepts either phrasing and holds both to naming +the model as the problem. """ from __future__ import annotations @@ -25,6 +28,8 @@ WEATHER_WAV = ( Path(__file__).resolve().parent / "realtime" / "fixtures" / "weather_question_24k.wav" ) +MISSING_MODEL_PHRASES: Final = ("model=none", "invalid model", "model is required") + class _OptionalTranscriptionForm(BaseModel): model: str | None = None @@ -105,8 +110,8 @@ class TestAudioTranscriptions: match result: case UnknownApiError(status_code=400, body=body): lowered: Final = body.lower() - assert "model" in lowered and ("required" in lowered or "invalid model" in lowered), ( - f"missing model error must identify the required model: {body[:300]}" + assert any(phrase in lowered for phrase in MISSING_MODEL_PHRASES), ( + f"missing model error must name the model as the problem: {body[:300]}" ) case other: pytest.fail(f"missing model expected a model-specific 400, got {other!r}") diff --git a/tests/e2e/models.py b/tests/e2e/models.py index ac41971a2c8..7711ca92b48 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -453,6 +453,7 @@ class AnthropicMessagesResponse(BaseModel): model: str | None = None content: list[AnthropicContentBlock] | None = None choices: list[ChatChoice] | None = None + usage: Usage | None = None class CountTokensResponse(BaseModel): @@ -717,8 +718,10 @@ class FineTuningJobsResponse(BaseModel): class LiteLLMParamsBody(BaseModel): """POST /model/new litellm_params: `model` is the only required field; `api_key` et al may be an `os.environ/FOO` reference the proxy resolves at call time. - `input_cost_per_token`/`output_cost_per_token` register a per-deployment custom - pricing override; left None (and dropped from the body) the deployment keeps the + The `*_cost_per_token` / `*_token_cost` fields register a per-deployment custom + pricing override (the cache and `_priority` rates only apply when both base + rates are set, which is what makes the proxy register the deployment's full + pricing entry); left None (and dropped from the body) the deployment keeps the backend's canonical rate.""" model: str @@ -745,6 +748,10 @@ class LiteLLMParamsBody(BaseModel): aws_external_id: str | None = None input_cost_per_token: float | None = None output_cost_per_token: float | None = None + cache_read_input_token_cost: float | None = None + cache_creation_input_token_cost: float | None = None + input_cost_per_token_priority: float | None = None + output_cost_per_token_priority: float | None = None extra_headers: dict[str, str] | None = None use_in_pass_through: bool | None = None complexity_router_config: dict[str, object] | None = None diff --git a/tests/e2e/provider_edge.py b/tests/e2e/provider_edge.py new file mode 100644 index 00000000000..ab0791e6b74 --- /dev/null +++ b/tests/e2e/provider_edge.py @@ -0,0 +1,546 @@ +"""Provider-edge record/replay server for e2e runs (LIT-5745). + +Record and replay scope to provider-bound traffic only: the proxy boots for +real, tests hit it for real, and only the hop from the proxy to the provider +is recorded or served from a bundle. Suites opt in per deployment by pointing +``litellm_params.api_base`` at ``provider_edge_api_base(mount)``, which is an +in-process HTTP server mounting each supported provider under a path prefix +(``http://127.0.0.1:/openai`` forwards to ``https://api.openai.com``). +In record mode the edge relays each request verbatim, stores the interaction, +and serves the proxy the same filtered response replay will serve later; in +replay mode it serves straight from the bundle and never opens a provider +connection, so a green replay run with a fake provider key proves the entire +proxy pipeline (auth, routing, spend logging) without provider spend. + +Request identity reuses fixture_canonical.py: interactions match by canonical +content key, order-independent across keys and FIFO within one. Edge requests +store no headers at all: SDK telemetry headers vary run to run and credential +headers must never touch disk. An unmatched replay call returns HTTP +``REPLAY_MISS_STATUS`` naming the closest recorded interaction, which the +proxy relays as a provider error the failing test surfaces. + +v1 limits: only the mounts in ``EDGE_MOUNTS`` (SigV4 providers like Bedrock +sign the Host header, so a forwarding edge breaks their signatures), JSON and +opaque single-part bodies (multipart boundaries are random per request), +streaming fidelity is LIT-5742, and CI wiring is LIT-5748. Suites that do not +wire the edge keep hitting providers live in every mode. +""" + +from __future__ import annotations + +import base64 +import difflib +import functools +import hashlib +import threading +from collections import deque +from collections.abc import Mapping +from dataclasses import dataclass, field +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from itertools import islice +from pathlib import Path +from types import MappingProxyType +from typing import Final, Literal, assert_never +from urllib.parse import parse_qsl, urlsplit + +from pydantic import JsonValue, TypeAdapter + +from e2e_http import NetworkError, RawResponse, forward +from fixture_bundle import ( + BundleRecorder, + Interaction, + LoadedBundle, + RecordedHttpResponse, + RecordedRequest, + UnreadableBundle, + UnsafeBundleDir, + interaction_filename, + load_bundle, + prepare_bundle, + slug_for_test, +) +from fixture_canonical import CanonicalRequest, canonical_string, canonicalize +from fixture_mode import ( + FIXTURE_MODES, + InvalidFixtureMode, + ReplayMiss, + current_test_key, + parse_fixture_mode, +) + +EDGE_MOUNTS: Final[Mapping[str, str]] = MappingProxyType( + { + "openai": "https://api.openai.com", + "anthropic": "https://api.anthropic.com", + } +) + +REPLAY_MISS_STATUS: Final = 599 + +_HOP_BY_HOP_HEADERS: Final[frozenset[str]] = frozenset( + { + "connection", + "keep-alive", + "proxy-authenticate", + "proxy-authorization", + "te", + "trailers", + "transfer-encoding", + "upgrade", + } +) +_REQUEST_DROPPED_HEADERS: Final[frozenset[str]] = _HOP_BY_HOP_HEADERS | { + "host", + "content-length", + "accept-encoding", +} +_RESPONSE_DROPPED_HEADERS: Final[frozenset[str]] = _HOP_BY_HOP_HEADERS | { + "content-encoding", + "content-length", + "set-cookie", +} + +_JSON: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) + + +def _edge_request(method: str, path: str, query: str, body: bytes | None) -> RecordedRequest: + """The identity replay matches on: the edge path (mount included), the query + as params, and the body as parsed JSON, or as a canonicalized content digest + when it is not JSON so opaque uploads still match across runs.""" + params: Final = dict(parse_qsl(query, keep_blank_values=True)) + if not body: + return RecordedRequest(method=method.lower(), path=path, headers={}, params=params) + decoded: Final = body.decode("utf-8", errors="replace") + try: + parsed: Final[JsonValue] = _JSON.validate_json(decoded) + except ValueError: + return RecordedRequest( + method=method.lower(), + path=path, + headers={}, + params=params, + file_sha256=hashlib.sha256(canonical_string(decoded).encode()).hexdigest(), + file_bytes=len(body), + ) + return RecordedRequest(method=method.lower(), path=path, headers={}, params=params, body=parsed) + + +def _build_pool(recorded: tuple[Interaction, ...]) -> dict[str, deque[Interaction]]: + keys: Final = tuple(canonicalize(interaction.request).key for interaction in recorded) + return { + key: deque( + interaction + for candidate_key, interaction in zip(keys, recorded, strict=True) + if candidate_key == key + ) + for key in dict.fromkeys(keys) + } + + +def _closest_recorded( + canonical: CanonicalRequest, recorded: tuple[Interaction, ...] +) -> tuple[CanonicalRequest, str]: + candidates: Final = tuple(canonicalize(interaction.request) for interaction in recorded) + ratios: Final = tuple( + difflib.SequenceMatcher( + None, f"{canonical.method} {canonical.path}\n{canonical.content}", + f"{candidate.method} {candidate.path}\n{candidate.content}", + ).ratio() + for candidate in candidates + ) + best: Final = max(range(len(candidates)), key=lambda index: ratios[index]) + return candidates[best], interaction_filename(best, recorded[best].request) + + +def _miss_message(test_key: str, slug: str, canonical: CanonicalRequest, bundle: LoadedBundle) -> str: + recorded: Final = bundle.interactions.get(slug, ()) + if not recorded: + return ( + f"replay miss for {test_key}: computed key {canonical.key} but nothing is recorded " + f"under {slug}; re-record with E2E_FIXTURE_MODE=record" + ) + closest, closest_file = _closest_recorded(canonical, recorded) + diff: Final = "\n".join( + islice( + difflib.unified_diff( + closest.pretty_content().splitlines(), + canonical.pretty_content().splitlines(), + fromfile=f"closest recorded ({closest_file})", + tofile="test made", + lineterm="", + ), + 60, + ) + ) + return ( + f"replay miss for {test_key}: no recorded interaction matches key {canonical.key}; " + f"closest recorded key is {closest.key} ({closest_file})\n{diff}\n" + "re-record with E2E_FIXTURE_MODE=record" + ) + + +@dataclass(slots=True) +class ReplaySource: + """One shared pool per test over a loaded bundle, so every provider call the + proxy makes in the session consumes from the same recorded interactions. + Every pool is built once at construction and per-key consumption is a single + atomic deque pop, so concurrent replay calls never race. Calls match by + canonical content key: order-independent across distinct keys (concurrent + tests interleave calls nondeterministically), FIFO within one key (a retry + or poll loop replays its recorded responses in recorded order).""" + + bundle: LoadedBundle + _pools: dict[str, dict[str, deque[Interaction]]] = field(init=False) + + def __post_init__(self) -> None: + self._pools = { + slug: _build_pool(recorded) for slug, recorded in self.bundle.interactions.items() + } + + def _pool(self, slug: str) -> dict[str, deque[Interaction]]: + return self._pools.get(slug, {}) + + def next_interaction(self, request: RecordedRequest) -> Interaction: + test_key: Final = current_test_key() + slug: Final = slug_for_test(test_key) + pool: Final = self._pool(slug) + canonical: Final = canonicalize(request) + queue: Final = pool.get(canonical.key) + if queue is None: + raise ReplayMiss(_miss_message(test_key, slug, canonical, self.bundle)) + try: + return queue.popleft() + except IndexError: + raise ReplayMiss( + f"replay exhausted for {test_key}: every recorded interaction for key " + f"{canonical.key} is already consumed; re-record with E2E_FIXTURE_MODE=record" + ) from None + + def leftover_error(self, test_key: str) -> str | None: + """Non-None when the test consumed fewer interactions than were recorded, + meaning a passing replay proved less than the bundle claims.""" + slug: Final = slug_for_test(test_key) + recorded: Final = self.bundle.interactions.get(slug, ()) + if not recorded: + return None + leftover: Final = tuple( + interaction for queue in self._pool(slug).values() for interaction in queue + ) + if not leftover: + return None + return ( + f"replay incomplete for {test_key}: {len(leftover)} of {len(recorded)} recorded " + f"interactions never consumed, e.g. {canonicalize(leftover[0].request).key}; " + "re-record with E2E_FIXTURE_MODE=record" + ) + + +@dataclass(frozen=True, slots=True) +class RecordEdge: + """Record backend: forward to the provider, persist, serve the filtered copy. + The lock serializes recorder writes because the edge server handles requests + on concurrent threads.""" + + recorder: BundleRecorder + lock: threading.Lock + + +@dataclass(frozen=True, slots=True) +class ReplayEdge: + source: ReplaySource + + +type EdgeBackend = RecordEdge | ReplayEdge + + +@dataclass(frozen=True, slots=True) +class EdgeReply: + status_code: int + headers: dict[str, str] + body: bytes + + +def _text_reply(status_code: int, message: str) -> EdgeReply: + return EdgeReply( + status_code=status_code, + headers={"content-type": "text/plain; charset=utf-8"}, + body=message.encode(), + ) + + +def _reply_from_recorded(response: RecordedHttpResponse) -> EdgeReply: + return EdgeReply( + status_code=response.status_code, + headers=dict(response.headers), + body=base64.b64decode(response.body_b64), + ) + + +def _recorded_response(outcome: RawResponse | NetworkError) -> RecordedHttpResponse: + match outcome: + case RawResponse(status_code=status_code, headers=headers, body=body): + return RecordedHttpResponse( + status_code=status_code, + headers={ + name: value + for name, value in headers.items() + if name not in _RESPONSE_DROPPED_HEADERS + }, + body_b64=base64.b64encode(body).decode("ascii"), + ) + case NetworkError(message=message): + return RecordedHttpResponse( + status_code=502, + headers={"content-type": "text/plain; charset=utf-8"}, + body_b64=base64.b64encode( + f"provider edge could not reach the provider: {message}".encode() + ).decode("ascii"), + ) + + +def _upstream_url(upstream_base: str, upstream_path: str, query: str) -> str: + url: Final = f"{upstream_base}/{upstream_path}" + return f"{url}?{query}" if query else url + + +def _handle_record( + backend: RecordEdge, + request: RecordedRequest, + *, + method: str, + url: str, + headers: Mapping[str, str], + body: bytes | None, + timeout: float, +) -> EdgeReply: + forwarded: Final = { + name: value for name, value in headers.items() if name.lower() not in _REQUEST_DROPPED_HEADERS + } + outcome: Final = forward(method, url, headers=forwarded, body=body, timeout=timeout) + response: Final = _recorded_response(outcome) + with backend.lock: + backend.recorder.record(test_key=current_test_key(), request=request, response=response) + return _reply_from_recorded(response) + + +def _handle_replay(source: ReplaySource, request: RecordedRequest) -> EdgeReply: + try: + interaction: Final = source.next_interaction(request) + except ReplayMiss as miss: + return _text_reply(REPLAY_MISS_STATUS, str(miss)) + return _reply_from_recorded(interaction.response) + + +def handle_edge_request( + backend: EdgeBackend, + mounts: Mapping[str, str], + method: str, + raw_path: str, + headers: Mapping[str, str], + body: bytes | None, + *, + timeout: float, +) -> EdgeReply: + """The edge's pure core, one HTTP exchange in and out: resolve the mount + prefix, then record (forward + persist) or replay (serve from the bundle). + Socket-free so unit tests exercise every branch without a server.""" + split: Final = urlsplit(raw_path) + mount, _, upstream_path = split.path.lstrip("/").partition("/") + upstream_base: Final = mounts.get(mount) + if upstream_base is None: + return _text_reply( + 404, f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(mounts))}" + ) + request: Final = _edge_request(method, split.path, split.query, body) + match backend: + case RecordEdge(): + return _handle_record( + backend, + request, + method=method, + url=_upstream_url(upstream_base, upstream_path, split.query), + headers=headers, + body=body, + timeout=timeout, + ) + case ReplayEdge(source=source): + return _handle_replay(source, request) + case _: + assert_never(backend) + + +class _EdgeHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_GET(self) -> None: + self._handle() + + def do_POST(self) -> None: + self._handle() + + def do_PUT(self) -> None: + self._handle() + + def do_PATCH(self) -> None: + self._handle() + + def do_DELETE(self) -> None: + self._handle() + + def _handle(self) -> None: + edge_server: Final = self.server + assert isinstance(edge_server, _EdgeHTTPServer) + length: Final = int(self.headers.get("content-length") or "0") + body: Final = self.rfile.read(length) if length else None + reply: Final = handle_edge_request( + edge_server.backend, + edge_server.mounts, + self.command, + self.path, + {name.lower(): value for name, value in self.headers.items()}, + body, + timeout=edge_server.forward_timeout, + ) + self.send_response(reply.status_code) + for name, value in reply.headers.items(): + self.send_header(name, value) + self.send_header("content-length", str(len(reply.body))) + self.end_headers() + self.wfile.write(reply.body) + + def log_message(self, format: str, *args: object) -> None: + """Silence the per-request stderr line BaseHTTPRequestHandler emits.""" + + +class _EdgeHTTPServer(ThreadingHTTPServer): + daemon_threads = True + + def __init__( + self, + bind: tuple[str, int], + *, + backend: EdgeBackend, + mounts: Mapping[str, str], + forward_timeout: float, + ) -> None: + super().__init__(bind, _EdgeHandler) + self.backend: Final = backend + self.mounts: Final = mounts + self.forward_timeout: Final = forward_timeout + + +@dataclass(frozen=True, slots=True) +class ProviderEdge: + port: int + advertise_host: str + + def api_base(self, mount: str) -> str: + return f"http://{self.advertise_host}:{self.port}/{mount}" + + +@dataclass(frozen=True, slots=True) +class RunningEdge: + edge: ProviderEdge + server: _EdgeHTTPServer + + def shutdown(self) -> None: + self.server.shutdown() + self.server.server_close() + + +def start_provider_edge( + backend: EdgeBackend, + *, + mounts: Mapping[str, str] = EDGE_MOUNTS, + bind_host: str = "127.0.0.1", + advertise_host: str | None = None, + forward_timeout: float = 60.0, +) -> RunningEdge: + """Boot an edge server on an OS-assigned port in a daemon thread. + ``advertise_host`` is what api_base URLs name (it differs from the bind + host when the proxy runs in a container and reaches the host machine via + a gateway address like host.docker.internal).""" + server: Final = _EdgeHTTPServer( + (bind_host, 0), backend=backend, mounts=mounts, forward_timeout=forward_timeout + ) + thread: Final = threading.Thread(target=server.serve_forever, name="e2e-provider-edge", daemon=True) + thread.start() + return RunningEdge( + edge=ProviderEdge(port=server.server_address[1], advertise_host=advertise_host or bind_host), + server=server, + ) + + +@functools.lru_cache(maxsize=8) +def _shared_recorder(root: Path) -> BundleRecorder: + prepared = prepare_bundle(root) + if isinstance(prepared, UnsafeBundleDir): + raise ValueError(f"E2E_FIXTURE_DIR {prepared.path} {prepared.reason}") + return prepared + + +@functools.lru_cache(maxsize=8) +def _shared_replay_source(root: Path) -> ReplaySource: + loaded = load_bundle(root) + if isinstance(loaded, UnreadableBundle): + raise ValueError(f"cannot replay from {root}: {loaded.reason}") + return ReplaySource(bundle=loaded) + + +@functools.lru_cache(maxsize=8) +def _shared_edge( + mode: Literal["record", "replay"], + bundle_dir: Path, + bind_host: str, + advertise_host: str, + forward_timeout: float, +) -> ProviderEdge: + backend: Final[EdgeBackend] = ( + RecordEdge(recorder=_shared_recorder(bundle_dir), lock=threading.Lock()) + if mode == "record" + else ReplayEdge(source=_shared_replay_source(bundle_dir)) + ) + return start_provider_edge( + backend, + mounts=EDGE_MOUNTS, + bind_host=bind_host, + advertise_host=advertise_host, + forward_timeout=forward_timeout, + ).edge + + +def replay_leftover_error(*, mode_raw: str, bundle_dir: Path, test_key: str) -> str | None: + """Teardown-time completeness check: in replay mode a passed test with + unconsumed recorded interactions must fail instead of passing against a + recording it no longer matches. Inert in every other mode.""" + if parse_fixture_mode(mode_raw) != "replay": + return None + return _shared_replay_source(bundle_dir).leftover_error(test_key) + + +def provider_edge_api_base( + mount: str, + *, + mode_raw: str, + bundle_dir: Path, + bind_host: str, + advertise_host: str, + forward_timeout: float = 60.0, +) -> str | None: + """The api_base a suite gives an edge-wired deployment: None in live mode + (the deployment keeps its real provider api_base) and the process-wide edge + server's mount URL in record and replay, booting the server on first use.""" + mode: Final = parse_fixture_mode(mode_raw) + match mode: + case InvalidFixtureMode(value=value): + raise ValueError(f"E2E_FIXTURE_MODE={value!r} is not one of {', '.join(FIXTURE_MODES)}") + case "live": + return None + case "record" | "replay": + if mount not in EDGE_MOUNTS: + raise ValueError( + f"unknown provider mount {mount!r}; known mounts: {', '.join(sorted(EDGE_MOUNTS))}" + ) + return _shared_edge(mode, bundle_dir, bind_host, advertise_host, forward_timeout).api_base(mount) + case _: + assert_never(mode) diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py index 3cae337a5ff..6cdd3354bf7 100644 --- a/tests/e2e/proxy_client.py +++ b/tests/e2e/proxy_client.py @@ -65,8 +65,6 @@ from models import ( ) from e2e_config import ( CONTROL_PLANE_BASE_URL, - FIXTURE_DIR, - FIXTURE_MODE_RAW, MASTER_KEY, POLL_INTERVAL, POLL_TIMEOUT, @@ -74,7 +72,6 @@ from e2e_config import ( REQUEST_TIMEOUT, settle_propagation, ) -from fixture_transport import select_transport from transport import HttpTransport, SplitTransport, Transport RowsPredicate = Callable[[list[SpendLogRow]], bool] @@ -547,9 +544,9 @@ def build_proxy_client( pass all three together, since a caller that overrides only the data plane would leave management calls pointed at the env default. - E2E_FIXTURE_MODE wraps (record) or replaces (replay) the transport here, so - every client built from this seam records or replays without changing shape; - unset it stays the plain SplitTransport (see fixture_transport.py).""" + Test-to-proxy traffic always goes over the wire, in every E2E_FIXTURE_MODE: + record and replay scope to the proxy's provider-bound calls via the + provider edge (see provider_edge.py), never to this transport.""" split = SplitTransport( data=HttpTransport( base_url=base_url, @@ -563,12 +560,7 @@ def build_proxy_client( ), ) return ProxyClient( - transport=select_transport( - split, - mode_raw=FIXTURE_MODE_RAW, - bundle_dir=FIXTURE_DIR, - master_key=master_key, - ), + transport=split, poll_timeout=POLL_TIMEOUT, poll_interval=POLL_INTERVAL, ) diff --git a/tests/e2e/quota_management/budgets/budget_client.py b/tests/e2e/quota_management/budgets/budget_client.py index 5b9253928af..543d5f959e0 100644 --- a/tests/e2e/quota_management/budgets/budget_client.py +++ b/tests/e2e/quota_management/budgets/budget_client.py @@ -460,7 +460,7 @@ class BudgetClient: time.sleep(_TEAM_READY_SLEEP_SECONDS) continue break - assert False, last_body + raise AssertionError(last_body) def update_team_member( self, diff --git a/tests/e2e/quota_management/spend_tracking/cost_rows.py b/tests/e2e/quota_management/spend_tracking/cost_rows.py new file mode 100644 index 00000000000..87af54fe83f --- /dev/null +++ b/tests/e2e/quota_management/spend_tracking/cost_rows.py @@ -0,0 +1,204 @@ +"""Cost-accounting helpers for the spend-tracking suite: the /spend/logs row shape +that carries the per-component cost breakdown, a poll that waits for it, and the +builders the cache-pricing tests share. + +The shared SpendLogRow deliberately stays thin (most tests only read totals), so +the component-cost tests model the metadata they assert on here instead: +`metadata.cost_breakdown` (input/output/cache-read/cache-creation/reasoning costs +plus the service-tier pricing basis) and `metadata.additional_usage_values` (the +cache token counts the biller derived from the provider's usage). + +Determinism strategy: every test registers its own deployment with explicit custom +rates for each component it asserts on (`register_priced_model`), so expected cost +is exactly tokens-on-the-row times configured rate, immune to provider price +changes. The rates are chosen ~100x above canonical and distinct from one another, +so a component billed at the wrong rate can never accidentally match. + +OpenAI prompt caching is implicit and keyed on the exact token prefix, with a +1024-token minimum. `cacheable_prefix` builds a prefix whose first word is the +run's unique marker: unique marker = the whole prefix is novel (a fresh cache +write), same marker + different question = a cache read that still misses the +proxy's own response cache. How long the prefix has to be before the provider +actually reports a read varies by model, so callers pass `words` to suit theirs. + +Two facts about the recorded bill that the assertions here encode, because the +two surfaces disagree on purpose. On the spend row, `input_cost` is gross: it +already contains the cache-read and cache-creation costs, so the row's total is +input + output + tool-usage and the fresh-token cost is input minus the two cache +components. In the response headers, `x-litellm-response-cost-input` is net of +cache, which is what makes the component headers sum to the total. +""" + +import time +from collections.abc import Callable + +from pydantic import BaseModel, RootModel + +from e2e_config import unique_marker +from e2e_http import Success +from lifecycle import ResourceManager +from models import LiteLLMParamsBody, SpendLogsParams +from proxy_client import ProxyClient + + +class CostBreakdownRow(BaseModel): + input_cost: float | None = None + output_cost: float | None = None + cache_read_cost: float | None = None + cache_creation_cost: float | None = None + reasoning_cost: float | None = None + tool_usage_cost: float | None = None + total_cost: float | None = None + service_tier: str | None = None + + +class AdditionalUsageValues(BaseModel): + cache_read_input_tokens: int | None = None + cache_creation_input_tokens: int | None = None + + +class CostRowMetadata(BaseModel): + cost_breakdown: CostBreakdownRow | None = None + additional_usage_values: AdditionalUsageValues | None = None + + +class CostRow(BaseModel): + request_id: str | None = None + spend: float | None = None + prompt_tokens: int | None = None + completion_tokens: int | None = None + metadata: CostRowMetadata | None = None + + @property + def breakdown(self) -> CostBreakdownRow: + assert self.metadata and self.metadata.cost_breakdown, ( + f"spend row {self.request_id} landed without a cost breakdown" + ) + return self.metadata.cost_breakdown + + @property + def cache_read_tokens(self) -> int: + if self.metadata and self.metadata.additional_usage_values: + return self.metadata.additional_usage_values.cache_read_input_tokens or 0 + return 0 + + @property + def cache_creation_tokens(self) -> int: + if self.metadata and self.metadata.additional_usage_values: + return self.metadata.additional_usage_values.cache_creation_input_tokens or 0 + return 0 + + +class CostRows(RootModel[list[CostRow]]): + pass + + +def approx_equal(actual: float, expected: float) -> bool: + """Within 1% or 1e-9 absolute - spend math, not exact float identity.""" + return abs(actual - expected) <= max(1e-9, abs(expected) * 1e-2) + + +def assert_total_is_sum_of_components(row: CostRow) -> None: + """The row's total is input + output + tool usage. The cache components are + already inside the gross input cost, so adding them again would double-bill.""" + breakdown = row.breakdown + components = sum( + cost or 0.0 + for cost in (breakdown.input_cost, breakdown.output_cost, breakdown.tool_usage_cost) + ) + assert breakdown.total_cost is not None and approx_equal(breakdown.total_cost, components), ( + f"total_cost {breakdown.total_cost} != input + output + tool usage ({components}): {breakdown}" + ) + assert row.spend is not None and approx_equal(row.spend, breakdown.total_cost), ( + f"row spend {row.spend} != breakdown total {breakdown.total_cost}" + ) + + +def assert_fresh_tokens_billed_at(row: CostRow, input_rate: float) -> None: + """Strip the cache components out of the gross input cost and what is left must + be the freshly-read tokens at the deployment's input rate.""" + breakdown = row.breakdown + fresh_tokens = (row.prompt_tokens or 0) - row.cache_read_tokens - row.cache_creation_tokens + fresh_cost = ( + (breakdown.input_cost or 0.0) + - (breakdown.cache_read_cost or 0.0) + - (breakdown.cache_creation_cost or 0.0) + ) + assert breakdown.input_cost is not None and approx_equal(fresh_cost, fresh_tokens * input_rate), ( + f"input_cost {breakdown.input_cost} less cache read {breakdown.cache_read_cost} and " + f"cache creation {breakdown.cache_creation_cost} leaves {fresh_cost}, not " + f"{fresh_tokens} fresh tokens * {input_rate} (prompt {row.prompt_tokens}, " + f"cache read {row.cache_read_tokens}, cache creation {row.cache_creation_tokens}); " + "cached tokens are being billed at the input rate" + ) + + +def poll_cost_row(proxy: ProxyClient, request_id: str) -> CostRow | None: + """Poll /spend/logs for the call's row until it lands with a cost breakdown + (rows flush ~60s behind the call via proxy_batch_write_at); None on timeout.""" + deadline = time.monotonic() + proxy.poll_timeout + while time.monotonic() < deadline: + result = proxy.transport.get( + "/spend/logs", + headers=proxy.transport.master, + params=SpendLogsParams(request_id=request_id), + response_type=CostRows, + ) + match result: + case Success(data=data): + rows = data.root + case _: + rows = [] + for row in rows: + if row.metadata and row.metadata.cost_breakdown: + return row + time.sleep(proxy.poll_interval) + return None + + +def poll_cost_row_where( + proxy: ProxyClient, api_key: str, predicate: Callable[[CostRow], bool] +) -> CostRow | None: + """Poll the key's own /spend/logs until one of its rows carries a cost breakdown + the predicate accepts; None on timeout. For calls whose response id is not the + id the bill is filed under, which is how a user finds the row in the UI anyway.""" + deadline = time.monotonic() + proxy.poll_timeout + while time.monotonic() < deadline: + result = proxy.transport.get( + "/spend/logs", + headers=proxy.transport.master, + params=SpendLogsParams(api_key=api_key), + response_type=CostRows, + ) + match result: + case Success(data=data): + rows = data.root + case _: + rows = [] + for row in rows: + if row.metadata and row.metadata.cost_breakdown and predicate(row): + return row + time.sleep(proxy.poll_interval) + return None + + +def register_priced_model( + proxy: ProxyClient, + resources: ResourceManager, + name_prefix: str, + litellm_params: LiteLLMParamsBody, +) -> str: + """Register a deployment with explicit custom rates (deleted on teardown) and + return its unique model name.""" + model_name = f"{name_prefix}-{unique_marker()}" + model_id = proxy.create_model(model_name, litellm_params) + resources.defer(lambda: proxy.delete_model(model_id)) + return model_name + + +def cacheable_prefix(marker: str, *, words: int = 1200) -> str: + """A prompt prefix above OpenAI's 1024-token caching minimum whose identity is + fully determined by `marker` (it is the first word, and prefix caching matches + from token zero). Raise `words` for models that only report a cache read on a + substantially longer prefix.""" + return " ".join(marker if i == 0 else f"token{i:04d}" for i in range(words)) diff --git a/tests/e2e/quota_management/spend_tracking/test_cache_cost_accounting_e2e.py b/tests/e2e/quota_management/spend_tracking/test_cache_cost_accounting_e2e.py new file mode 100644 index 00000000000..c50ec3d902f --- /dev/null +++ b/tests/e2e/quota_management/spend_tracking/test_cache_cost_accounting_e2e.py @@ -0,0 +1,287 @@ +"""Live e2e: prompt-cache token accounting bills each cache component at its own rate. + +Four regressions the gateway has shipped fixes for, pinned against real OpenAI +prompt caching (implicit, keyed on the token prefix). Every test registers its own +deployment with distinct custom rates for input / output / cache-read / +cache-creation, so the expected bill is exactly the row's token counts times the +configured rates and a component billed at the wrong rate can never pass: + +- cache writes: gpt-5.6's cache-write tokens must land on the spend row as + cache-creation tokens billed at the cache-creation rate, not silently at the + input rate (#34046) +- breakdown components: the row's metadata.cost_breakdown must itemize cache-read, + cache-creation, and reasoning costs, with reasoning a subset of output (#31686) +- streaming: a streamed call's reassembled usage must keep the cached-token detail + so cache reads bill at the cache-read discount, not full input price (#34812) +- /v1/messages bridge: a request served by a Responses-only OpenAI model crosses + the anthropic-messages -> Responses adapter and must keep its cache-read tokens + and their discounted billing (#34957) + +Each test drives the model that actually reports the component it bills, which is +not the same model throughout. gpt-5.6-luna reports cache-write tokens on every +call over the caching minimum and never reports a cache read, so it is the one +model that can prove cache-write billing and the one model that can never prove +cache-read billing. gpt-5.5 is the reverse: it reports cached tokens on the second +call and no cache writes at all. gpt-5.3-codex is Responses-only, which is what +forces the /v1/messages bridge, and it starts reporting cache reads once the +prefix is a few thousand tokens rather than one. + +OpenAI caching is best-effort, so each test retries with a fresh prefix (new +marker = brand-new cache identity) up to three times before failing; the prime and +measured calls share the prefix but differ in the trailing question, which defeats +the proxy's own response cache without touching the provider's prefix cache. + +The test that asserts on reasoning cost requests reasoning explicitly with +`reasoning_effort`, so that assertion rests on a parameter the test sets rather +than on whatever the model happens to do by default. Its prime call carries the +same value: OpenAI's prefix cache keys on the reasoning setting as well as the +tokens, so a prime at a different effort never produces a read. +""" + +import pytest + +from cost_rows import ( + CostRow, + approx_equal, + assert_fresh_tokens_billed_at, + assert_total_is_sum_of_components, + cacheable_prefix, + poll_cost_row, + poll_cost_row_where, + register_priced_model, +) +from e2e_config import unique_marker +from e2e_http import unwrap +from lifecycle import ResourceManager +from models import AnthropicMessagesBody, ChatBody, ChatMessage, LiteLLMParamsBody +from pydantic import BaseModel +from spend_e2e_client import SpendClient + +pytestmark = pytest.mark.e2e + +CACHE_WRITE_BACKEND = "openai/gpt-5.6-luna" +CACHE_READ_BACKEND = "openai/gpt-5.5" +BRIDGE_BACKEND = "openai/gpt-5.3-codex" +BRIDGE_PREFIX_WORDS = 3000 +OPENAI_API_KEY = "os.environ/OPENAI_API_KEY" +CACHE_ATTEMPTS = 3 + +INPUT_RATE = 4e-05 +OUTPUT_RATE = 8e-05 +CACHE_READ_RATE = 1e-05 +CACHE_WRITE_RATE = 5e-05 + +PRIME_QUESTION = "Reply with the single word ready." +REASONING_QUESTION = "Compute 47*83 - 19*7 step by step, then reply with just the final number." +REASONING_EFFORT = "high" + + +class _StreamChunk(BaseModel): + id: str | None = None + + +def _cache_priced_params(backend: str) -> LiteLLMParamsBody: + return LiteLLMParamsBody( + model=backend, + api_key=OPENAI_API_KEY, + input_cost_per_token=INPUT_RATE, + output_cost_per_token=OUTPUT_RATE, + cache_read_input_token_cost=CACHE_READ_RATE, + cache_creation_input_token_cost=CACHE_WRITE_RATE, + ) + + +def _chat_body( + model: str, content: str, *, stream: bool = False, reasoning_effort: str | None = None +) -> ChatBody: + return ChatBody( + model=model, + messages=[ChatMessage(role="user", content=content)], + stream=stream, + max_completion_tokens=4000, + reasoning_effort=reasoning_effort, + ) + + +def _require_row(client: SpendClient, request_id: str) -> CostRow: + row = poll_cost_row(client.proxy, request_id) + assert row is not None, f"no spend row with a cost breakdown landed for {request_id}" + return row + + +def _assert_cache_read_billed(row: CostRow) -> None: + assert row.breakdown.cache_read_cost is not None and approx_equal( + row.breakdown.cache_read_cost, row.cache_read_tokens * CACHE_READ_RATE + ), ( + f"cache_read_cost {row.breakdown.cache_read_cost} != " + f"{row.cache_read_tokens} cached tokens * {CACHE_READ_RATE}" + ) + assert_fresh_tokens_billed_at(row, INPUT_RATE) + assert_total_is_sum_of_components(row) + + +class TestCacheCostAccounting: + @pytest.mark.covers("quota_management.spend_tracking.cache_write.bills_cache_creation_rate") + def test_cache_write_tokens_billed_at_cache_creation_rate( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, resources, "cache-write-priced", _cache_priced_params(CACHE_WRITE_BACKEND) + ) + + for _ in range(CACHE_ATTEMPTS): + prompt = f"{cacheable_prefix(unique_marker())}\n{PRIME_QUESTION}" + chat = unwrap(client.proxy.chat(scoped_key, _chat_body(model, prompt))) + assert chat.id, f"chat response carried no id: {chat}" + row = _require_row(client, chat.id) + if row.cache_creation_tokens > 0: + break + else: + pytest.fail( + f"OpenAI reported no cache-write tokens across {CACHE_ATTEMPTS} fresh " + "~2k-token prompts; the cache-write billing path was never exercised" + ) + + assert row.breakdown.cache_creation_cost is not None and approx_equal( + row.breakdown.cache_creation_cost, row.cache_creation_tokens * CACHE_WRITE_RATE + ), ( + f"cache_creation_cost {row.breakdown.cache_creation_cost} != " + f"{row.cache_creation_tokens} cache-write tokens * {CACHE_WRITE_RATE}" + ) + assert_fresh_tokens_billed_at(row, INPUT_RATE) + assert_total_is_sum_of_components(row) + + @pytest.mark.covers("quota_management.spend_tracking.cost_breakdown.reports_component_costs") + def test_cost_breakdown_reports_component_costs( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, resources, "breakdown-priced", _cache_priced_params(CACHE_READ_BACKEND) + ) + + for _ in range(CACHE_ATTEMPTS): + prefix = cacheable_prefix(unique_marker()) + unwrap( + client.proxy.chat( + scoped_key, + _chat_body( + model, f"{prefix}\n{PRIME_QUESTION}", reasoning_effort=REASONING_EFFORT + ), + ) + ) + chat = unwrap( + client.proxy.chat( + scoped_key, + _chat_body( + model, + f"{prefix}\n{REASONING_QUESTION}", + reasoning_effort=REASONING_EFFORT, + ), + ) + ) + assert chat.id, f"chat response carried no id: {chat}" + row = _require_row(client, chat.id) + if row.cache_read_tokens > 0: + break + else: + pytest.fail( + f"no cache read landed across {CACHE_ATTEMPTS} prime+read rounds; " + "the component-cost breakdown was never exercised with cached input" + ) + + usage = chat.usage + assert usage is not None and usage.completion_tokens_details is not None, ( + f"no completion token details on the measured call: {chat}" + ) + reasoning_tokens = usage.completion_tokens_details.reasoning_tokens or 0 + assert reasoning_tokens > 0, f"the reasoning question produced no reasoning tokens: {usage}" + + breakdown = row.breakdown + assert breakdown.output_cost is not None and approx_equal( + breakdown.output_cost, (row.completion_tokens or 0) * OUTPUT_RATE + ), ( + f"output_cost {breakdown.output_cost} != " + f"{row.completion_tokens} completion tokens * {OUTPUT_RATE}" + ) + assert breakdown.reasoning_cost is not None and approx_equal( + breakdown.reasoning_cost, reasoning_tokens * OUTPUT_RATE + ), ( + f"reasoning_cost {breakdown.reasoning_cost} != " + f"{reasoning_tokens} reasoning tokens * {OUTPUT_RATE}" + ) + assert breakdown.reasoning_cost <= (breakdown.output_cost or 0.0) * 1.01, ( + f"reasoning_cost {breakdown.reasoning_cost} exceeds output_cost " + f"{breakdown.output_cost}; reasoning must be a subset of output" + ) + _assert_cache_read_billed(row) + + @pytest.mark.covers("quota_management.spend_tracking.stream_cache_read.bills_cache_read_rate") + def test_streaming_cache_read_billed_at_cache_read_rate( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, resources, "stream-cache-priced", _cache_priced_params(CACHE_READ_BACKEND) + ) + + for _ in range(CACHE_ATTEMPTS): + prefix = cacheable_prefix(unique_marker()) + unwrap(client.proxy.chat(scoped_key, _chat_body(model, f"{prefix}\n{PRIME_QUESTION}"))) + result = client.proxy.chat_stream( + scoped_key, + _chat_body(model, f"{prefix}\nReply with the single word cached.", stream=True), + ) + assert result.ok and result.stream_events, ( + f"streamed chat failed (status {result.status_code}): {result.body[:300]}" + ) + stream_id = _StreamChunk.model_validate_json(result.stream_events[0]).id + assert stream_id, f"first stream chunk carried no id: {result.stream_events[0][:200]}" + row = _require_row(client, stream_id) + if row.cache_read_tokens > 0: + break + else: + pytest.fail( + f"no cache read landed across {CACHE_ATTEMPTS} prime+stream rounds; " + "streaming cache-read billing was never exercised" + ) + + _assert_cache_read_billed(row) + + @pytest.mark.covers("quota_management.spend_tracking.messages_bridge.keeps_cache_tokens") + def test_messages_bridge_keeps_cache_tokens( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, resources, "bridge-cache-priced", _cache_priced_params(BRIDGE_BACKEND) + ) + + def bridge_call(content: str) -> int: + response = unwrap( + client.proxy.messages( + scoped_key, + AnthropicMessagesBody( + model=model, + messages=[ChatMessage(role="user", content=content)], + max_tokens=4000, + ), + ) + ) + assert response.usage is not None, f"bridged response carried no usage: {response}" + return response.usage.cache_read_input_tokens or 0 + + for _ in range(CACHE_ATTEMPTS): + prefix = cacheable_prefix(unique_marker(), words=BRIDGE_PREFIX_WORDS) + bridge_call(f"{prefix}\n{PRIME_QUESTION}") + if bridge_call(f"{prefix}\nReply with the single word bridged.") > 0: + break + else: + pytest.fail( + f"no cache read survived {CACHE_ATTEMPTS} bridged prime+read rounds; " + "cache tokens are not surviving the anthropic-messages -> Responses bridge" + ) + + row = poll_cost_row_where(client.proxy, scoped_key, lambda r: r.cache_read_tokens > 0) + assert row is not None, ( + "the bridged call reported cached tokens but no spend row for the key " + "recorded any; the cache tokens were dropped on the way to the bill" + ) + _assert_cache_read_billed(row) diff --git a/tests/e2e/quota_management/spend_tracking/test_cost_headers_e2e.py b/tests/e2e/quota_management/spend_tracking/test_cost_headers_e2e.py new file mode 100644 index 00000000000..203be611905 --- /dev/null +++ b/tests/e2e/quota_management/spend_tracking/test_cost_headers_e2e.py @@ -0,0 +1,136 @@ +"""Live e2e: the per-component x-litellm-response-cost-* headers keep their contract. + +Pins the header contract shipped in #36965: alongside the x-litellm-response-cost +total, every response carries the component costs (input, output, cache-read, +cache-creation, reasoning, tool-usage), where input covers only fresh tokens (the +cache components are subtracted out) so the components sum to the total, and +reasoning stays a subset of output. + +The deployment carries distinct custom rates per component, a prime call fills the +provider's prefix cache, and the measured call re-reads it, so the cache-read +header is exercised with a real nonzero value instead of passing vacuously. The +backend is gpt-5.5 because it reports cached tokens on the second call; the +gpt-5.6 line reports cache writes and never a read, which would leave the +cache-read header at zero forever. The raw-transport send is used because the +typed chat client validates bodies and drops headers. OpenAI caching is +best-effort, so the prime+measure round retries with a fresh prefix before +failing. +""" + +import pytest + +from cost_rows import approx_equal, cacheable_prefix, register_priced_model +from e2e_config import unique_marker +from e2e_http import StreamingResponse +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, ChatResponse, LiteLLMParamsBody +from spend_e2e_client import SpendClient + +pytestmark = pytest.mark.e2e + +BACKEND = "openai/gpt-5.5" +OPENAI_API_KEY = "os.environ/OPENAI_API_KEY" +CACHE_ATTEMPTS = 3 + +INPUT_RATE = 4e-05 +OUTPUT_RATE = 8e-05 +CACHE_READ_RATE = 1e-05 +CACHE_WRITE_RATE = 5e-05 + +COMPONENT_HEADERS = ( + "x-litellm-response-cost-input", + "x-litellm-response-cost-cache-read", + "x-litellm-response-cost-cache-creation", + "x-litellm-response-cost-output", + "x-litellm-response-cost-tool-usage", +) + + +def _header_cost(response: StreamingResponse, name: str) -> float: + value = response.headers.get(name) + return float(value) if value not in (None, "", "None") else 0.0 + + +class TestCostHeaders: + @pytest.mark.covers("quota_management.spend_tracking.cost_headers.additive_components") + def test_component_cost_headers_sum_to_total( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, + resources, + "header-priced", + LiteLLMParamsBody( + model=BACKEND, + api_key=OPENAI_API_KEY, + input_cost_per_token=INPUT_RATE, + output_cost_per_token=OUTPUT_RATE, + cache_read_input_token_cost=CACHE_READ_RATE, + cache_creation_input_token_cost=CACHE_WRITE_RATE, + ), + ) + + def priced_call(content: str) -> StreamingResponse: + response = client.proxy.transport.send( + "/chat/completions", + headers=client.proxy.transport.bearer(scoped_key), + json=ChatBody( + model=model, + messages=[ChatMessage(role="user", content=content)], + max_completion_tokens=4000, + ), + ) + assert response.ok, f"chat failed (status {response.status_code}): {response.body[:300]}" + return response + + for _ in range(CACHE_ATTEMPTS): + prefix = cacheable_prefix(unique_marker()) + priced_call(f"{prefix}\nReply with the single word ready.") + measured = priced_call(f"{prefix}\nReply with the single word measured.") + if _header_cost(measured, "x-litellm-response-cost-cache-read") > 0: + break + else: + pytest.fail( + f"no cache read landed across {CACHE_ATTEMPTS} prime+measure rounds; " + "the cache-read cost header was never exercised with a nonzero value" + ) + + total = measured.response_cost + assert total is not None and total > 0, ( + f"x-litellm-response-cost missing or zero: {measured.headers}" + ) + component_sum = sum(_header_cost(measured, name) for name in COMPONENT_HEADERS) + assert approx_equal(component_sum, total), ( + f"component headers sum to {component_sum}, not the total {total}: " + f"{ {name: measured.headers.get(name) for name in COMPONENT_HEADERS} }" + ) + + reasoning = _header_cost(measured, "x-litellm-response-cost-reasoning") + output = _header_cost(measured, "x-litellm-response-cost-output") + assert reasoning <= output * 1.01, ( + f"reasoning header {reasoning} exceeds output header {output}; " + "reasoning must be a subset of output" + ) + + usage = ChatResponse.model_validate_json(measured.body).usage + assert usage is not None, f"measured response carried no usage: {measured.body[:300]}" + cached_tokens = ( + usage.prompt_tokens_details.cached_tokens or 0 if usage.prompt_tokens_details else 0 + ) + cache_creation_tokens = usage.cache_creation_input_tokens or 0 + assert cached_tokens > 0, f"cache-read header nonzero but usage shows no cached tokens: {usage}" + assert approx_equal( + _header_cost(measured, "x-litellm-response-cost-cache-read"), + cached_tokens * CACHE_READ_RATE, + ), ( + f"cache-read header {measured.headers.get('x-litellm-response-cost-cache-read')} != " + f"{cached_tokens} cached tokens * {CACHE_READ_RATE}" + ) + fresh_tokens = (usage.prompt_tokens or 0) - cached_tokens - cache_creation_tokens + assert approx_equal( + _header_cost(measured, "x-litellm-response-cost-input"), fresh_tokens * INPUT_RATE + ), ( + f"input header {measured.headers.get('x-litellm-response-cost-input')} != " + f"{fresh_tokens} fresh tokens * {INPUT_RATE}; the input component is not " + "subtracting the cache components" + ) diff --git a/tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py b/tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py new file mode 100644 index 00000000000..ced7c819d42 --- /dev/null +++ b/tests/e2e/quota_management/spend_tracking/test_provider_edge_spend_e2e.py @@ -0,0 +1,50 @@ +"""The provider-edge demonstrator: one spend-tracking flow wired through the +record/replay edge (LIT-5745). + +This is the reference for wiring a suite to the edge: register a deployment +whose ``api_base`` comes from ``e2e_config.provider_edge_base``, then exercise +the proxy exactly as a live test would. In live mode the base is None and the +deployment talks to the real provider; in record mode it talks through the +local edge, which forwards to the provider and captures the exchange; in +replay mode the same test drives the REAL proxy and REAL database on the +recorded provider traffic alone, so key auth, routing, and the spend-log +write path are all still under test with zero provider calls. +""" + +import pytest + +from e2e_config import CHEAP_OPENAI_MODEL, provider_edge_base +from lifecycle import ResourceManager +from models import LiteLLMParamsBody +from spend_e2e_client import SpendClient, unique_marker, unwrap + +pytestmark = pytest.mark.e2e + + +@pytest.mark.covers("quota_management.spend_tracking.chat_completions.logs_cost") +def test_edge_wired_chat_writes_nonzero_spend_row( + client: SpendClient, resources: ResourceManager, scoped_key: str +) -> None: + base = provider_edge_base("openai") + model = f"e2e-edge-openai-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody( + model=f"openai/{CHEAP_OPENAI_MODEL}", + api_key="os.environ/OPENAI_API_KEY", + api_base=None if base is None else f"{base}/v1", + ), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + chat = unwrap( + client.chat(scoped_key, model, f"reply with one word {unique_marker()}", max_tokens=16) + ) + assert chat.id + + rows = client.poll_logs_for_key( + scoped_key, predicate=lambda rs: any((r.spend or 0) > 0 for r in rs) + ) + matching = [row for row in rows if row.request_id == chat.id] + assert matching, f"no SpendLogs row for request_id {chat.id}; saw {len(rows)} row(s)" + assert (matching[0].spend or 0) > 0, f"spend row for {chat.id} has zero spend" diff --git a/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py b/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py new file mode 100644 index 00000000000..770c5699b4e --- /dev/null +++ b/tests/e2e/quota_management/spend_tracking/test_service_tier_pricing_e2e.py @@ -0,0 +1,121 @@ +"""Live e2e: a service_tier request bills every component at the tier's own rates. + +Pins the tier-billing fixes (#35923, #35925): a priority-tier call must price +input and output at the deployment's `*_priority` rates, including the reasoning +tokens inside output (the shipped bug billed reasoning at the default-tier rate), +and the spend row must record the tier the bill was computed on. + +The deployment carries custom base AND priority rates, each distinct, so a bill +computed from the wrong tier (or a mix) cannot match the expected numbers. The +prompt is a fresh unique marker per run, keeping cached tokens out of the math. +The response's own `service_tier` echo is asserted first: if OpenAI ever declined +priority processing and served the default tier, the test fails there instead of +producing a vacuous rate comparison. Reasoning is requested explicitly with +`reasoning_effort`, so the reasoning-rate assertion rests on a parameter the test +sets rather than on whatever the model happens to do by default. +""" + +import pytest + +from cost_rows import ( + approx_equal, + assert_fresh_tokens_billed_at, + assert_total_is_sum_of_components, + poll_cost_row, + register_priced_model, +) +from e2e_config import unique_marker +from e2e_http import unwrap +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, LiteLLMParamsBody +from spend_e2e_client import SpendClient + +pytestmark = pytest.mark.e2e + +BACKEND = "openai/gpt-5.6-luna" +OPENAI_API_KEY = "os.environ/OPENAI_API_KEY" + +INPUT_RATE = 4e-05 +OUTPUT_RATE = 8e-05 +PRIORITY_INPUT_RATE = 6e-05 +PRIORITY_OUTPUT_RATE = 1.6e-04 + +REASONING_EFFORT = "high" + + +class TestServiceTierPricing: + @pytest.mark.covers("quota_management.spend_tracking.service_tier.bills_tier_rates") + def test_priority_tier_bills_priority_rates( + self, client: SpendClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = register_priced_model( + client.proxy, + resources, + "tier-priced", + LiteLLMParamsBody( + model=BACKEND, + api_key=OPENAI_API_KEY, + input_cost_per_token=INPUT_RATE, + output_cost_per_token=OUTPUT_RATE, + input_cost_per_token_priority=PRIORITY_INPUT_RATE, + output_cost_per_token_priority=PRIORITY_OUTPUT_RATE, + ), + ) + + chat = unwrap( + client.proxy.chat( + scoped_key, + ChatBody( + model=model, + messages=[ + ChatMessage( + role="user", + content=( + f"{unique_marker()} Compute 47*83 - 19*7 step by step, " + "then reply with just the final number." + ), + ) + ], + max_completion_tokens=4000, + service_tier="priority", + reasoning_effort=REASONING_EFFORT, + ), + ) + ) + assert chat.service_tier == "priority", ( + f"OpenAI served tier {chat.service_tier!r} instead of priority; " + "tier billing was never exercised" + ) + assert chat.id, f"chat response carried no id: {chat}" + + row = poll_cost_row(client.proxy, chat.id) + assert row is not None, f"no spend row with a cost breakdown landed for {chat.id}" + breakdown = row.breakdown + + assert breakdown.service_tier == "priority", ( + f"the bill records pricing basis {breakdown.service_tier!r}, not priority" + ) + + assert_fresh_tokens_billed_at(row, PRIORITY_INPUT_RATE) + assert breakdown.output_cost is not None and approx_equal( + breakdown.output_cost, (row.completion_tokens or 0) * PRIORITY_OUTPUT_RATE + ), ( + f"output_cost {breakdown.output_cost} != {row.completion_tokens} tokens * priority rate " + f"{PRIORITY_OUTPUT_RATE} (base rate would give {(row.completion_tokens or 0) * OUTPUT_RATE})" + ) + + usage = chat.usage + assert usage is not None and usage.completion_tokens_details is not None, ( + f"no completion token details on the priority call: {chat}" + ) + reasoning_tokens = usage.completion_tokens_details.reasoning_tokens or 0 + assert reasoning_tokens > 0, f"the reasoning question produced no reasoning tokens: {usage}" + assert breakdown.reasoning_cost is not None and approx_equal( + breakdown.reasoning_cost, reasoning_tokens * PRIORITY_OUTPUT_RATE + ), ( + f"reasoning_cost {breakdown.reasoning_cost} != {reasoning_tokens} reasoning tokens * " + f"priority rate {PRIORITY_OUTPUT_RATE} (the default-tier rate would give " + f"{reasoning_tokens * OUTPUT_RATE})" + ) + + assert_total_is_sum_of_components(row) diff --git a/tests/e2e/test_fixture_bundle.py b/tests/e2e/test_fixture_bundle.py index fd4cca6451f..b49ab565e39 100644 --- a/tests/e2e/test_fixture_bundle.py +++ b/tests/e2e/test_fixture_bundle.py @@ -1,9 +1,9 @@ -"""Harness coverage for the on-disk fixture bundle format (LIT-5729). +"""Harness coverage for the on-disk fixture bundle format (LIT-5729/LIT-5745). No proxy and no ``e2e`` marker: these pin the bundle CONTRACT - the seven-day freshness gate that names the bundle's age, record mode's wipe safety (never delete a directory that is not a bundle), collision-free per-test slugs, and -lossless Result round-trips - so replay can never silently drift from what +grouped-in-order loading - so replay can never silently drift from what record wrote. """ @@ -12,18 +12,6 @@ from __future__ import annotations from datetime import datetime, timedelta, timezone from pathlib import Path -import pytest -from pydantic import BaseModel - -from e2e_http import ( - NetworkError, - RateLimitedError, - Result, - Success, - UnauthorizedError, - UnknownApiError, - ValidationError, -) from fixture_bundle import ( BUNDLE_FORMAT_VERSION, MANIFEST_FILENAME, @@ -32,28 +20,22 @@ from fixture_bundle import ( FreshBundle, LoadedBundle, Manifest, + RecordedHttpResponse, RecordedRequest, - RecordedResult, StaleBundle, UnreadableBundle, UnsafeBundleDir, check_freshness, format_age, - from_result, interaction_filename, load_bundle, prepare_bundle, slug_for_test, - to_result, ) NOW = datetime(2026, 8, 18, 12, 0, 0, tzinfo=timezone.utc) -class Payload(BaseModel): - value: str - - def write_manifest( root: Path, recorded_at: datetime, *, format_version: int = BUNDLE_FORMAT_VERSION ) -> None: @@ -74,20 +56,8 @@ def plain_request(path: str) -> RecordedRequest: return RecordedRequest(method="post", path=path, headers={}) -class TestResultRoundTrip: - @pytest.mark.parametrize( - "result", - [ - Success(status_code=201, data=Payload(value="ok")), - NetworkError(message="connection refused"), - UnauthorizedError(), - RateLimitedError(retry_after_seconds=7, body="slow down"), - ValidationError(message="bad shape"), - UnknownApiError(status_code=502, body="upstream exploded"), - ], - ) - def test_every_result_kind_survives_disk_and_back(self, result: Result[Payload]) -> None: - assert to_result(from_result(result), Payload) == result +def plain_response() -> RecordedHttpResponse: + return RecordedHttpResponse(status_code=401, headers={}, body_b64="") class TestFreshness: @@ -144,7 +114,7 @@ class TestPrepareBundle: prepared(root).record( test_key="old.py::test_old", request=plain_request("/stale"), - response=RecordedResult(kind="unauthorized"), + response=plain_response(), ) assert any(entry.is_dir() for entry in root.iterdir()) prepared(root) @@ -193,7 +163,7 @@ class TestRecordAndLoad: recorder.record( test_key=key, request=plain_request(path), - response=RecordedResult(kind="unauthorized"), + response=plain_response(), ) loaded = load_bundle(root) assert isinstance(loaded, LoadedBundle) @@ -208,7 +178,7 @@ class TestRecordAndLoad: recorder.record( test_key=key, request=plain_request(f"/{key[-3:]}"), - response=RecordedResult(kind="unauthorized"), + response=plain_response(), ) loaded = load_bundle(root) assert isinstance(loaded, LoadedBundle) diff --git a/tests/e2e/test_fixture_canonical.py b/tests/e2e/test_fixture_canonical.py index 30c57dc3ac6..8890848522c 100644 --- a/tests/e2e/test_fixture_canonical.py +++ b/tests/e2e/test_fixture_canonical.py @@ -140,11 +140,21 @@ class TestKeyDistinctness: second = request(headers={"traceparent": "00-cc-dd-01", "x-api-key": "two"}) assert canonicalize(first).key == canonicalize(second).key + def test_query_params_are_identity(self) -> None: + first = request("get", "/v1/vector_stores", params={"limit": "100"}) + second = request("get", "/v1/vector_stores", params={"limit": "10"}) + assert canonicalize(first).key != canonicalize(second).key + def test_secret_set_versus_unset_stays_distinct(self) -> None: with_key = request(body={"api_key": "sk-live-aaaaaaaaaaaaaaaa"}) without_key = request(body={"api_key": None}) assert canonicalize(with_key).key != canonicalize(without_key).key + def test_form_fields_are_identity(self) -> None: + first = request("upload", "/v1/files", form={"purpose": "assistants"}, file_sha256="a" * 64) + second = request("upload", "/v1/files", form={"purpose": "batch"}, file_sha256="a" * 64) + assert canonicalize(first).key != canonicalize(second).key + def test_file_content_is_identity(self) -> None: first = request( "upload", "/v1/files", file_name="batch.jsonl", file_sha256="a" * 64, file_bytes=10 diff --git a/tests/e2e/test_fixture_mode.py b/tests/e2e/test_fixture_mode.py new file mode 100644 index 00000000000..109bb9e1b11 --- /dev/null +++ b/tests/e2e/test_fixture_mode.py @@ -0,0 +1,114 @@ +"""Harness coverage for fixture-mode selection and determinism (LIT-5729/LIT-5745). + +No proxy and no ``e2e`` marker. Pins the mode parser, the deterministic +per-test marker sequence a replay run must regenerate, the collection-time +gate (including the stale message that names the bundle's age), and the pytest +report header. The provider-edge record/replay behavior itself is pinned in +test_provider_edge.py. +""" + +from __future__ import annotations + +import hashlib +from datetime import datetime, timedelta, timezone +from pathlib import Path + +import pytest + +from fixture_bundle import BUNDLE_FORMAT_VERSION, MANIFEST_FILENAME, Manifest +from fixture_mode import ( + InvalidFixtureMode, + current_test_key, + deterministic_marker, + fixture_mode_collection_error, + fixture_report_lines, + parse_fixture_mode, +) + +NOW = datetime(2026, 8, 18, 12, 0, 0, tzinfo=timezone.utc) + + +def write_manifest(root: Path, recorded_at: datetime) -> None: + root.mkdir(parents=True, exist_ok=True) + manifest = Manifest( + format_version=BUNDLE_FORMAT_VERSION, recorded_at=recorded_at, harness_version="abc1234" + ) + (root / MANIFEST_FILENAME).write_text(manifest.model_dump_json(), encoding="utf-8") + + +class TestParseFixtureMode: + @pytest.mark.parametrize( + ("raw", "expected"), + [("live", "live"), ("record", "record"), ("replay", "replay"), ("", "live"), (" REPLAY ", "replay")], + ) + def test_known_values_normalize(self, raw: str, expected: str) -> None: + assert parse_fixture_mode(raw) == expected + + def test_unknown_value_is_invalid_with_the_original_spelling(self) -> None: + assert parse_fixture_mode("cached") == InvalidFixtureMode(value="cached") + + +class TestDeterministicMarker: + def test_sequence_is_a_pure_function_of_test_and_ordinal(self) -> None: + """A replay process must regenerate exactly the markers the record + process generated, so the Nth marker of a test is pinned to a pure + function of the node id and N.""" + key = current_test_key() + assert deterministic_marker() == hashlib.sha1(f"{key}#0".encode()).hexdigest()[:12] + assert deterministic_marker() == hashlib.sha1(f"{key}#1".encode()).hexdigest()[:12] + + +class TestCurrentTestKey: + def test_names_this_test_and_strips_the_phase(self) -> None: + key = current_test_key() + assert key.endswith("TestCurrentTestKey::test_names_this_test_and_strips_the_phase") + assert "(call)" not in key + + +class TestCollectionGate: + def test_invalid_mode_names_the_value_and_the_choices(self, tmp_path: Path) -> None: + assert ( + fixture_mode_collection_error("cached", tmp_path, now=NOW) + == "E2E_FIXTURE_MODE='cached' is not one of live, record, replay" + ) + + @pytest.mark.parametrize("mode_raw", ["live", "", "record"]) + def test_live_and_record_never_block_collection(self, mode_raw: str, tmp_path: Path) -> None: + assert fixture_mode_collection_error(mode_raw, tmp_path / "missing", now=NOW) is None + + def test_replay_with_no_bundle_says_how_to_record_one(self, tmp_path: Path) -> None: + reason = fixture_mode_collection_error("replay", tmp_path / "missing", now=NOW) + assert reason is not None + assert f"no {MANIFEST_FILENAME}" in reason + assert "E2E_FIXTURE_MODE=record" in reason + + def test_stale_replay_bundle_fails_naming_its_age(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + write_manifest(root, NOW - timedelta(days=9, hours=5)) + reason = fixture_mode_collection_error("replay", root, now=NOW) + assert reason is not None + assert "age 9d5h exceeds the 7-day limit" in reason + assert "re-record with E2E_FIXTURE_MODE=record" in reason + + def test_fresh_replay_bundle_collects(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + write_manifest(root, NOW - timedelta(days=2)) + assert fixture_mode_collection_error("replay", root, now=NOW) is None + + +class TestReportHeader: + def test_live_mode_prints_nothing(self, tmp_path: Path) -> None: + assert fixture_report_lines("live", tmp_path, now=NOW) == [] + assert fixture_report_lines("", tmp_path, now=NOW) == [] + + def test_record_and_replay_name_the_bundle(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + recorded_at = NOW - timedelta(days=1) + write_manifest(root, recorded_at) + assert fixture_report_lines("record", root, now=NOW) == [ + f"e2e fixture mode: record -> {root}" + ] + replay_lines = fixture_report_lines("replay", root, now=NOW) + assert len(replay_lines) == 1 + assert "replay" in replay_lines[0] + assert recorded_at.isoformat() in replay_lines[0] diff --git a/tests/e2e/test_fixture_transport.py b/tests/e2e/test_fixture_transport.py deleted file mode 100644 index e61088d841c..00000000000 --- a/tests/e2e/test_fixture_transport.py +++ /dev/null @@ -1,676 +0,0 @@ -"""Harness coverage for the record/replay transports (LIT-5729). - -No proxy and no ``e2e`` marker. A fake in-memory ``Transport`` stands in for -the live one (dependency injection, no monkeypatching): recording must pass -every value through unchanged while writing one redacted interaction file per -call, and replay must serve identical values from the bundle alone - the -fake's call log proves nothing reaches the inner transport - failing hard -(``ReplayMiss``) on any content drift, printing the computed canonical key and -the closest recorded key (LIT-5741; the pure canonicalizer is pinned in -test_fixture_canonical.py). The collection-time gate and report header are -pinned here too, including the stale message that names the bundle's age. -""" - -from __future__ import annotations - -import hashlib -import sys -import threading -from concurrent.futures import ThreadPoolExecutor -from dataclasses import dataclass, field -from datetime import datetime, timedelta, timezone -from pathlib import Path -from uuid import uuid4 - -import pytest -from pydantic import BaseModel - -from e2e_http import ( - AuthHeaders, - BinaryStream, - ProbeResult, - Result, - StreamingResponse, - Success, -) -from fixture_bundle import ( - BUNDLE_FORMAT_VERSION, - MANIFEST_FILENAME, - BundleRecorder, - Interaction, - LoadedBundle, - Manifest, - RecordedResult, - load_bundle, - prepare_bundle, - slug_for_test, -) -from fixture_canonical import canonicalize -from fixture_transport import ( - InvalidFixtureMode, - RecordingTransport, - ReplayMiss, - ReplaySource, - ReplayTransport, - current_test_key, - deterministic_marker, - fixture_mode_collection_error, - fixture_report_lines, - parse_fixture_mode, - recorded_request, - replay_leftover_error, - select_transport, -) -from transport import Transport - -NOW = datetime(2026, 8, 18, 12, 0, 0, tzinfo=timezone.utc) - - -class Payload(BaseModel): - value: str - - -class Body(BaseModel): - prompt: str - - -class Query(BaseModel): - q: str - - -class DeployParams(BaseModel): - model: str - api_key: str | None = None - aws_secret_access_key: str | None = None - - -class DeployBody(BaseModel): - model_name: str - litellm_params: DeployParams - - -STREAMING = StreamingResponse( - status_code=200, - body="", - content_type="text/event-stream", - chunks=2, - stream_events=["one", "two"], - stream_done=True, -) -BINARY = BinaryStream(status_code=200, content_type="audio/mpeg", chunk_count=3, total_bytes=42) -PROBE = ProbeResult(status_code=200, body="alive") - - -@dataclass -class FakeTransport: - calls: list[str] = field(default_factory=list) - - def bearer(self, key: str) -> AuthHeaders: - return AuthHeaders(authorization=f"Bearer {key}") - - @property - def master(self) -> AuthHeaders: - return self.bearer("sk-fake-master") - - def _success[R: BaseModel](self, response_type: type[R]) -> Result[R]: - return Success(status_code=200, data=response_type.model_validate({"value": "live"})) - - def post[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - self.calls.append(f"post {path}") - return self._success(response_type) - - def get[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - params: BaseModel, - response_type: type[R], - timeout: float | None = None, - ) -> Result[R]: - self.calls.append(f"get {path}") - return self._success(response_type) - - def delete[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - json: BaseModel, - response_type: type[R], - params: BaseModel | None = None, - ) -> Result[R]: - self.calls.append(f"delete {path}") - return self._success(response_type) - - def patch[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - self.calls.append(f"patch {path}") - return self._success(response_type) - - def put[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] - ) -> Result[R]: - self.calls.append(f"put {path}") - return self._success(response_type) - - def stream(self, path: str, *, headers: BaseModel, json: BaseModel) -> StreamingResponse: - self.calls.append(f"stream {path}") - return STREAMING - - def stream_binary( - self, path: str, *, headers: BaseModel, json: BaseModel, chunk_size: int = 8192 - ) -> BinaryStream: - self.calls.append(f"stream_binary {path}") - return BINARY - - def send( - self, - path: str, - *, - headers: BaseModel, - json: BaseModel, - params: BaseModel | None = None, - stream: bool = False, - ) -> StreamingResponse: - self.calls.append(f"send {path}") - return STREAMING - - def probe(self, path: str, *, params: BaseModel) -> ProbeResult: - self.calls.append(f"probe {path}") - return PROBE - - def upload[R: BaseModel]( - self, - path: str, - *, - headers: BaseModel, - form: BaseModel, - filename: str, - content: bytes, - file_content_type: str = "application/jsonl", - file_field: str = "file", - params: BaseModel | None = None, - response_type: type[R], - ) -> Result[R]: - self.calls.append(f"upload {path}") - return self._success(response_type) - - def download(self, path: str, *, headers: BaseModel) -> StreamingResponse: - self.calls.append(f"download {path}") - return STREAMING - - -def make_recorder(root: Path) -> BundleRecorder: - recorder = prepare_bundle(root) - assert isinstance(recorder, BundleRecorder) - return recorder - - -def replay_source(root: Path) -> ReplaySource: - loaded = load_bundle(root) - assert isinstance(loaded, LoadedBundle) - return ReplaySource(bundle=loaded) - - -def this_tests_files(root: Path) -> list[Path]: - slug_dir = root / slug_for_test(current_test_key()) - return sorted(slug_dir.glob("*.json")) if slug_dir.is_dir() else [] - - -def write_manifest(root: Path, recorded_at: datetime) -> None: - root.mkdir(parents=True, exist_ok=True) - manifest = Manifest( - format_version=BUNDLE_FORMAT_VERSION, recorded_at=recorded_at, harness_version="abc1234" - ) - (root / MANIFEST_FILENAME).write_text(manifest.model_dump_json(), encoding="utf-8") - - -class TestParseFixtureMode: - @pytest.mark.parametrize( - ("raw", "expected"), - [("live", "live"), ("record", "record"), ("replay", "replay"), ("", "live"), (" REPLAY ", "replay")], - ) - def test_known_values_normalize(self, raw: str, expected: str) -> None: - assert parse_fixture_mode(raw) == expected - - def test_unknown_value_is_invalid_with_the_original_spelling(self) -> None: - assert parse_fixture_mode("cached") == InvalidFixtureMode(value="cached") - - -class TestDeterministicMarker: - def test_sequence_is_a_pure_function_of_test_and_ordinal(self) -> None: - """A replay process must regenerate exactly the markers the record - process generated, so the Nth marker of a test is pinned to a pure - function of the node id and N.""" - key = current_test_key() - assert deterministic_marker() == hashlib.sha1(f"{key}#0".encode()).hexdigest()[:12] - assert deterministic_marker() == hashlib.sha1(f"{key}#1".encode()).hexdigest()[:12] - - -class TestCurrentTestKey: - def test_names_this_test_and_strips_the_phase(self) -> None: - key = current_test_key() - assert key.endswith("TestCurrentTestKey::test_names_this_test_and_strips_the_phase") - assert "(call)" not in key - - -class TestRecordingTransport: - def test_passes_the_result_through_and_writes_one_file_per_call(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - result = recording.post( - "/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload - ) - assert result == Success(status_code=200, data=Payload(value="live")) - assert fake.calls == ["post /model/new"] - files = this_tests_files(root) - assert [file.name for file in files] == ["0000-post-model-new.json"] - interaction = Interaction.model_validate_json(files[0].read_text(encoding="utf-8")) - assert interaction.request.method == "post" - assert interaction.request.path == "/model/new" - - def test_redacts_auth_header_values_in_the_recorded_request(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - headers = AuthHeaders.model_validate( - {"authorization": "Bearer sk-secret", "x-litellm-api-key": "sk-other"} - ) - recording.post("/key/generate", headers=headers, json=Body(prompt="x"), response_type=Payload) - interaction = Interaction.model_validate_json( - this_tests_files(root)[0].read_text(encoding="utf-8") - ) - assert interaction.request.headers == { - "authorization": "", - "x-litellm-api-key": "", - } - assert "sk-secret" not in this_tests_files(root)[0].read_text(encoding="utf-8") - - def test_redacts_credential_body_fields_in_the_recorded_request(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.post( - "/model/new", - headers=fake.master, - json=DeployBody( - model_name="m", - litellm_params=DeployParams(model="openai/gpt", api_key="sk-live-provider-secret-123456"), - ), - response_type=Payload, - ) - raw = this_tests_files(root)[0].read_text(encoding="utf-8") - interaction = Interaction.model_validate_json(raw) - assert "sk-live-provider-secret-123456" not in raw - assert isinstance(interaction.request.body, dict) - params = interaction.request.body["litellm_params"] - assert isinstance(params, dict) - assert params["api_key"] == "" - assert params["aws_secret_access_key"] is None - - def test_upload_records_a_content_digest_not_the_bytes(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.upload( - "/v1/files", - headers=fake.master, - form=Query(q="batch"), - filename="batch.jsonl", - content=b'{"custom_id": "1"}', - response_type=Payload, - ) - interaction = Interaction.model_validate_json( - this_tests_files(root)[0].read_text(encoding="utf-8") - ) - assert interaction.request.file_name == "batch.jsonl" - assert interaction.request.file_bytes == len(b'{"custom_id": "1"}') - assert interaction.request.file_sha256 is not None - assert "custom_id" not in interaction.request.model_dump_json() - - -class TestReplayTransport: - def test_serves_recorded_values_without_touching_the_inner_transport( - self, tmp_path: Path - ) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recorded_post = recording.post( - "/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload - ) - recorded_get = recording.get( - "/v1/models", headers=fake.master, params=Query(q="all"), response_type=Payload - ) - recorded_stream = recording.stream( - "/chat/completions", headers=fake.master, json=Body(prompt="hi") - ) - recorded_probe = recording.probe("/health/liveliness", params=Query(q="1")) - recorded_binary = recording.stream_binary( - "/v1/audio/speech", headers=fake.master, json=Body(prompt="say") - ) - calls_after_record = list(fake.calls) - - replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234") - assert ( - replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload) - == recorded_post - ) - assert ( - replay.get("/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload) - == recorded_get - ) - assert ( - replay.stream("/chat/completions", headers=replay.master, json=Body(prompt="hi")) - == recorded_stream - ) - assert replay.probe("/health/liveliness", params=Query(q="1")) == recorded_probe - assert ( - replay.stream_binary("/v1/audio/speech", headers=replay.master, json=Body(prompt="say")) - == recorded_binary - ) - assert fake.calls == calls_after_record - - def test_miss_names_the_computed_key_and_the_closest_recorded_key(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload) - replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234") - with pytest.raises(ReplayMiss) as excinfo: - replay.get("/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload) - message = str(excinfo.value) - assert "no recorded interaction matches key get /v1/models #" in message - assert "closest recorded key is post /model/new #" in message - assert "0000-post-model-new.json" in message - assert "re-record with E2E_FIXTURE_MODE=record" in message - - def test_content_drift_on_the_same_route_misses_with_no_live_call(self, tmp_path: Path) -> None: - """The naive verb+path match replayed a stale response for a request - whose content had changed, silently passing; a content key must miss, - print both canonical forms' diff, and never reach the inner transport.""" - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload) - calls_after_record = list(fake.calls) - replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234") - with pytest.raises(ReplayMiss) as excinfo: - replay.post("/model/new", headers=replay.master, json=Body(prompt="y"), response_type=Payload) - message = str(excinfo.value) - assert "no recorded interaction matches key post /model/new #" in message - assert "closest recorded key is post /model/new #" in message - assert '- "prompt": "x"' in message - assert '+ "prompt": "y"' in message - assert fake.calls == calls_after_record - - def test_exhausted_key_names_the_key(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload) - replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234") - replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload) - with pytest.raises( - ReplayMiss, match=r"every recorded interaction for key post /model/new #\w{16} is already consumed" - ): - replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload) - - def test_replays_out_of_recorded_order_across_distinct_keys(self, tmp_path: Path) -> None: - """Concurrent tests interleave independent calls nondeterministically - (e.g. a burst of parallel chat calls), so replay matches by content, - never by recorded position.""" - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload) - recording.post("/key/generate", headers=fake.master, json=Body(prompt="k"), response_type=Payload) - source = replay_source(root) - replay: Transport = ReplayTransport(source=source, master_key="sk-1234") - replay.post("/key/generate", headers=replay.master, json=Body(prompt="k"), response_type=Payload) - replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload) - assert source.leftover_error(current_test_key()) is None - - def test_identical_requests_replay_their_responses_in_recorded_order(self, tmp_path: Path) -> None: - """A poll loop makes the same request repeatedly and asserts on the - progression, so duplicates under one key stay FIFO.""" - root = tmp_path / "bundle" - recorder = make_recorder(root) - recorder.record( - test_key=current_test_key(), - request=recorded_request( - "get", "/v1/models", headers=AuthHeaders(authorization="Bearer sk-x"), params=Query(q="all") - ), - response=RecordedResult(kind="success", status_code=200, data={"value": "first"}), - ) - recorder.record( - test_key=current_test_key(), - request=recorded_request( - "get", "/v1/models", headers=AuthHeaders(authorization="Bearer sk-x"), params=Query(q="all") - ), - response=RecordedResult(kind="success", status_code=200, data={"value": "second"}), - ) - replay: Transport = ReplayTransport(source=replay_source(root), master_key="sk-1234") - first = replay.get("/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload) - second = replay.get("/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload) - assert first == Success(status_code=200, data=Payload(value="first")) - assert second == Success(status_code=200, data=Payload(value="second")) - - def test_concurrent_replays_of_one_key_serve_each_recording_exactly_once(self, tmp_path: Path) -> None: - """A burst of parallel identical calls consumes one shared pool: no - response duplicated, none forgotten, nothing left over at teardown. - The tiny switch interval forces thread preemption inside pool setup - and consumption, so a non-atomic pool build or pop fails this test.""" - root = tmp_path / "bundle" - recorder = make_recorder(root) - for ordinal in range(32): - recorder.record( - test_key=current_test_key(), - request=recorded_request( - "get", "/v1/models", headers=AuthHeaders(authorization="Bearer sk-x"), params=Query(q="all") - ), - response=RecordedResult(kind="success", status_code=200, data={"value": f"v{ordinal:02d}"}), - ) - source = replay_source(root) - replay: Transport = ReplayTransport(source=source, master_key="sk-1234") - barrier = threading.Barrier(8) - - def consume_one() -> str: - result = replay.get( - "/v1/models", headers=replay.master, params=Query(q="all"), response_type=Payload - ) - assert isinstance(result, Success) - return result.data.value - - def consume(_: int) -> tuple[str, ...]: - barrier.wait() - return tuple(consume_one() for _call in range(4)) - - previous_interval = sys.getswitchinterval() - sys.setswitchinterval(1e-6) - try: - with ThreadPoolExecutor(max_workers=8) as executor: - served = sorted(value for values in executor.map(consume, range(8)) for value in values) - finally: - sys.setswitchinterval(previous_interval) - assert served == [f"v{ordinal:02d}" for ordinal in range(32)] - assert source.leftover_error(current_test_key()) is None - - -class TestRecordedKeySets: - def test_two_separate_recordings_of_one_flow_produce_identical_key_sets( - self, tmp_path: Path - ) -> None: - """Everything a run randomizes (markers, virtual keys, dates) must - canonicalize out, so separately recorded runs of the same suite agree - on every match key and a bundle recorded elsewhere replays here.""" - - def record_flow(root: Path, run_date: str) -> list[str]: - fake = FakeTransport() - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - marker = deterministic_marker() - recording.post( - "/model/new", - headers=fake.master, - json=DeployBody( - model_name=f"e2e-chat-{marker}", - litellm_params=DeployParams(model="openai/gpt", api_key=f"sk-live-{uuid4().hex}"), - ), - response_type=Payload, - ) - recording.post( - "/chat/completions", - headers=recording.bearer(f"sk-{uuid4().hex}"), - json=Body(prompt=f"Reply with the single word ok. {marker}"), - response_type=Payload, - ) - recording.get( - "/spend/logs", headers=fake.master, params=Query(q=run_date), response_type=Payload - ) - loaded = load_bundle(root) - assert isinstance(loaded, LoadedBundle) - return sorted( - canonicalize(interaction.request).key - for interactions in loaded.interactions.values() - for interaction in interactions - ) - - first_keys = record_flow(tmp_path / "one", "2026-08-18") - second_keys = record_flow(tmp_path / "two", "2026-08-19") - assert first_keys == second_keys - assert len(first_keys) == 3 - - -class TestReplayLeftover: - def test_fully_consumed_recording_leaves_nothing(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload) - source = replay_source(root) - replay: Transport = ReplayTransport(source=source, master_key="sk-1234") - replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload) - assert source.leftover_error(current_test_key()) is None - - def test_unconsumed_trailing_interactions_name_the_next_call(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload) - recording.probe("/health/liveliness", params=Query(q="1")) - source = replay_source(root) - replay: Transport = ReplayTransport(source=source, master_key="sk-1234") - replay.post("/model/new", headers=replay.master, json=Body(prompt="x"), response_type=Payload) - error = source.leftover_error(current_test_key()) - assert error is not None - assert "1 of 2 recorded interactions never consumed" in error - assert "e.g. probe /health/liveliness #" in error - assert "re-record with E2E_FIXTURE_MODE=record" in error - - def test_test_without_recordings_has_no_leftover(self, tmp_path: Path) -> None: - root = tmp_path / "bundle" - make_recorder(root) - assert replay_source(root).leftover_error("suite.py::test_never_recorded") is None - - def test_inert_outside_replay_mode(self, tmp_path: Path) -> None: - missing = tmp_path / "missing" - assert replay_leftover_error(mode_raw="", bundle_dir=missing, test_key="k") is None - assert replay_leftover_error(mode_raw="record", bundle_dir=missing, test_key="k") is None - - def test_replay_mode_reads_the_shared_bundle(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - recording: Transport = RecordingTransport(inner=fake, recorder=make_recorder(root)) - recording.post("/model/new", headers=fake.master, json=Body(prompt="x"), response_type=Payload) - error = replay_leftover_error(mode_raw="replay", bundle_dir=root, test_key=current_test_key()) - assert error is not None - assert "1 of 1 recorded interactions never consumed" in error - - -class TestSelectTransport: - def test_live_returns_the_live_transport_untouched(self, tmp_path: Path) -> None: - fake = FakeTransport() - for mode_raw in ("live", ""): - assert ( - select_transport(fake, mode_raw=mode_raw, bundle_dir=tmp_path / "b", master_key="sk") - is fake - ) - - def test_record_wraps_live_and_starts_a_fresh_bundle(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - write_manifest(root, NOW - timedelta(days=30)) - (root / "old-test-slug").mkdir() - (root / "old-test-slug" / "0000-post-old.json").write_text("{}", encoding="utf-8") - selected = select_transport(fake, mode_raw="record", bundle_dir=root, master_key="sk") - assert isinstance(selected, RecordingTransport) - assert selected.inner is fake - assert {entry.name for entry in root.iterdir()} == {MANIFEST_FILENAME} - - def test_replay_builds_a_transport_from_the_bundle_alone(self, tmp_path: Path) -> None: - fake = FakeTransport() - root = tmp_path / "bundle" - make_recorder(root) - selected = select_transport(fake, mode_raw="replay", bundle_dir=root, master_key="sk-master") - assert isinstance(selected, ReplayTransport) - assert selected.master == AuthHeaders(authorization="Bearer sk-master") - - def test_invalid_mode_raises_naming_the_value(self, tmp_path: Path) -> None: - with pytest.raises(ValueError, match="cached"): - select_transport( - FakeTransport(), mode_raw="cached", bundle_dir=tmp_path / "b", master_key="sk" - ) - - -class TestCollectionGate: - def test_invalid_mode_names_the_value_and_the_choices(self, tmp_path: Path) -> None: - assert ( - fixture_mode_collection_error("cached", tmp_path, now=NOW) - == "E2E_FIXTURE_MODE='cached' is not one of live, record, replay" - ) - - @pytest.mark.parametrize("mode_raw", ["live", "", "record"]) - def test_live_and_record_never_block_collection(self, mode_raw: str, tmp_path: Path) -> None: - assert fixture_mode_collection_error(mode_raw, tmp_path / "missing", now=NOW) is None - - def test_replay_with_no_bundle_says_how_to_record_one(self, tmp_path: Path) -> None: - reason = fixture_mode_collection_error("replay", tmp_path / "missing", now=NOW) - assert reason is not None - assert f"no {MANIFEST_FILENAME}" in reason - assert "E2E_FIXTURE_MODE=record" in reason - - def test_stale_replay_bundle_fails_naming_its_age(self, tmp_path: Path) -> None: - root = tmp_path / "bundle" - write_manifest(root, NOW - timedelta(days=9, hours=5)) - reason = fixture_mode_collection_error("replay", root, now=NOW) - assert reason is not None - assert "age 9d5h exceeds the 7-day limit" in reason - assert "re-record with E2E_FIXTURE_MODE=record" in reason - - def test_fresh_replay_bundle_collects(self, tmp_path: Path) -> None: - root = tmp_path / "bundle" - write_manifest(root, NOW - timedelta(days=2)) - assert fixture_mode_collection_error("replay", root, now=NOW) is None - - -class TestReportHeader: - def test_live_mode_prints_nothing(self, tmp_path: Path) -> None: - assert fixture_report_lines("live", tmp_path, now=NOW) == [] - assert fixture_report_lines("", tmp_path, now=NOW) == [] - - def test_record_and_replay_name_the_bundle(self, tmp_path: Path) -> None: - root = tmp_path / "bundle" - recorded_at = NOW - timedelta(days=1) - write_manifest(root, recorded_at) - assert fixture_report_lines("record", root, now=NOW) == [ - f"e2e fixture mode: record -> {root}" - ] - replay_lines = fixture_report_lines("replay", root, now=NOW) - assert len(replay_lines) == 1 - assert "replay" in replay_lines[0] - assert recorded_at.isoformat() in replay_lines[0] diff --git a/tests/e2e/test_provider_edge.py b/tests/e2e/test_provider_edge.py new file mode 100644 index 00000000000..492eee57aaf --- /dev/null +++ b/tests/e2e/test_provider_edge.py @@ -0,0 +1,492 @@ +"""Harness coverage for the provider-edge record/replay server (LIT-5745). + +No proxy and no ``e2e`` marker. A stdlib http.server stands in for the +provider (dependency injection via the mounts mapping, no monkeypatching): +record mode must forward each edge call to it verbatim, persist one +interaction file, and serve the proxy the same filtered response replay will +serve later; replay mode must serve byte-identical responses from the bundle +alone, with the fake provider's hit log proving nothing leaves the process, +and answer any drifted call with HTTP ``REPLAY_MISS_STATUS`` naming the +computed and closest recorded canonical keys (LIT-5741; the pure canonicalizer +is pinned in test_fixture_canonical.py). Requests are made through +``e2e_http.forward`` so the whole HTTP surface of the edge is exercised; the +pure ``handle_edge_request`` core is pinned socket-free alongside. +""" + +from __future__ import annotations + +import base64 +import json +import threading +from collections.abc import Generator, Mapping +from concurrent.futures import ThreadPoolExecutor +from contextlib import contextmanager +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from pathlib import Path + +import pytest +from pydantic import TypeAdapter + +from e2e_http import RawResponse, forward +from fixture_bundle import ( + BundleRecorder, + Interaction, + LoadedBundle, + RecordedHttpResponse, + RecordedRequest, + load_bundle, + prepare_bundle, + slug_for_test, +) +from fixture_mode import current_test_key +from provider_edge import ( + REPLAY_MISS_STATUS, + EdgeBackend, + ProviderEdge, + RecordEdge, + ReplayEdge, + ReplaySource, + handle_edge_request, + provider_edge_api_base, + replay_leftover_error, + start_provider_edge, +) + +CHAT_PATH = "/openai/v1/chat/completions" +REPLAY_MOUNTS = {"openai": "https://replay.invalid"} +JSON_OBJECT = TypeAdapter(dict[str, object]) + + +def json_object(body: bytes) -> dict[str, object]: + return JSON_OBJECT.validate_json(body) + + +class _FakeProvider(ThreadingHTTPServer): + daemon_threads = True + + def __init__(self, bind: tuple[str, int]) -> None: + super().__init__(bind, _FakeProviderHandler) + self.hits: list[str] = [] + + +class _FakeProviderHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + + def do_POST(self) -> None: + self._respond() + + def do_GET(self) -> None: + self._respond() + + def _respond(self) -> None: + provider = self.server + assert isinstance(provider, _FakeProvider) + length = int(self.headers.get("content-length") or "0") + body = self.rfile.read(length) if length else b"" + provider.hits.append(f"{self.command} {self.path}") + payload = json.dumps( + {"echo": body.decode("utf-8"), "path": self.path, "hit": len(provider.hits)} + ).encode() + self.send_response(200) + self.send_header("content-type", "application/json") + self.send_header("content-length", str(len(payload))) + self.send_header("x-upstream", "fake") + self.send_header("set-cookie", "session=fake-cookie") + self.end_headers() + self.wfile.write(payload) + + def log_message(self, format: str, *args: object) -> None: + """Silence the per-request stderr line BaseHTTPRequestHandler emits.""" + + +@contextmanager +def fake_provider() -> Generator[_FakeProvider]: + server = _FakeProvider(("127.0.0.1", 0)) + thread = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + try: + yield server + finally: + server.shutdown() + server.server_close() + + +def provider_url(server: _FakeProvider) -> str: + return f"http://127.0.0.1:{server.server_address[1]}" + + +@contextmanager +def running_edge(backend: EdgeBackend, mounts: Mapping[str, str]) -> Generator[ProviderEdge]: + running = start_provider_edge(backend, mounts=mounts, bind_host="127.0.0.1") + try: + yield running.edge + finally: + running.shutdown() + + +def record_backend(root: Path) -> RecordEdge: + recorder = prepare_bundle(root) + assert isinstance(recorder, BundleRecorder) + return RecordEdge(recorder=recorder, lock=threading.Lock()) + + +def replay_source(root: Path) -> ReplaySource: + loaded = load_bundle(root) + assert isinstance(loaded, LoadedBundle) + return ReplaySource(bundle=loaded) + + +def call_edge( + edge: ProviderEdge, + method: str, + path: str, + *, + body: bytes | None = None, + headers: dict[str, str] | None = None, +) -> RawResponse: + outcome = forward( + method, + f"http://{edge.advertise_host}:{edge.port}{path}", + headers=headers or {}, + body=body, + timeout=10.0, + ) + assert isinstance(outcome, RawResponse) + return outcome + + +def this_tests_files(root: Path) -> list[Path]: + slug_dir = root / slug_for_test(current_test_key()) + return sorted(slug_dir.glob("*.json")) if slug_dir.is_dir() else [] + + +def chat_body(prompt: str) -> bytes: + return json.dumps({"model": "gpt", "messages": [{"role": "user", "content": prompt}]}).encode() + + +class TestRecordMode: + def test_forwards_to_the_provider_and_writes_one_interaction_file(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + reply = call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + assert provider.hits == ["POST /v1/chat/completions"] + assert reply.status_code == 200 + served = json_object(reply.body) + assert served["echo"] == chat_body("hi").decode() + files = this_tests_files(root) + assert [file.name for file in files] == ["0000-post-openai-v1-chat-completions.json"] + interaction = Interaction.model_validate_json(files[0].read_text(encoding="utf-8")) + assert interaction.request.method == "post" + assert interaction.request.path == CHAT_PATH + assert interaction.request.body == json_object(chat_body("hi")) + assert interaction.response.status_code == 200 + + def test_never_stores_headers_so_credentials_never_touch_disk(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge( + edge, + "POST", + CHAT_PATH, + body=chat_body("hi"), + headers={"authorization": "Bearer sk-live-provider-secret-abc123"}, + ) + raw = this_tests_files(root)[0].read_text(encoding="utf-8") + assert "sk-live-provider-secret-abc123" not in raw + interaction = Interaction.model_validate_json(raw) + assert interaction.request.headers == {} + + def test_strips_volatile_response_headers_and_serves_the_filtered_copy(self, tmp_path: Path) -> None: + """What record serves the proxy must equal what replay will serve later + (record/replay parity), so the filtered stored copy is served in both.""" + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + reply = call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + assert reply.headers.get("x-upstream") == "fake" + assert "set-cookie" not in reply.headers + interaction = Interaction.model_validate_json( + this_tests_files(root)[0].read_text(encoding="utf-8") + ) + assert interaction.response.headers.get("x-upstream") == "fake" + assert "set-cookie" not in interaction.response.headers + assert "content-length" not in interaction.response.headers + + def test_unreachable_provider_records_and_serves_a_502(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with running_edge(record_backend(root), {"openai": "http://127.0.0.1:9"}) as edge: + reply = call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + assert reply.status_code == 502 + assert b"could not reach the provider" in reply.body + interaction = Interaction.model_validate_json( + this_tests_files(root)[0].read_text(encoding="utf-8") + ) + assert interaction.response.status_code == 502 + + +class TestReplayMode: + def test_serves_recorded_bytes_with_zero_provider_hits(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + recorded = call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + hits_after_record = list(provider.hits) + with running_edge( + ReplayEdge(source=replay_source(root)), {"openai": provider_url(provider)} + ) as edge: + replayed = call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + assert provider.hits == hits_after_record + assert replayed.status_code == recorded.status_code + assert replayed.body == recorded.body + assert replayed.headers.get("x-upstream") == "fake" + + def test_request_identity_ignores_auth_headers(self, tmp_path: Path) -> None: + """The proxy sends different bearer tokens across runs (fresh virtual + keys, rotated provider keys), so headers are no part of the match.""" + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge( + edge, "POST", CHAT_PATH, body=chat_body("hi"), + headers={"authorization": "Bearer sk-first-run"}, + ) + with running_edge(ReplayEdge(source=replay_source(root)), REPLAY_MOUNTS) as edge: + replayed = call_edge( + edge, "POST", CHAT_PATH, body=chat_body("hi"), + headers={"authorization": "Bearer sk-second-run"}, + ) + assert replayed.status_code == 200 + + def test_content_drift_returns_the_miss_status_naming_both_keys(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge(edge, "POST", CHAT_PATH, body=chat_body("x")) + with running_edge(ReplayEdge(source=replay_source(root)), REPLAY_MOUNTS) as edge: + missed = call_edge(edge, "POST", CHAT_PATH, body=chat_body("y")) + assert missed.status_code == REPLAY_MISS_STATUS + message = missed.body.decode() + assert f"no recorded interaction matches key post {CHAT_PATH} #" in message + assert f"closest recorded key is post {CHAT_PATH} #" in message + assert '"content": "x"' in message + assert '"content": "y"' in message + assert "re-record with E2E_FIXTURE_MODE=record" in message + + def test_query_params_are_part_of_the_identity(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge(edge, "GET", "/openai/v1/models?purpose=batch") + assert provider.hits == ["GET /v1/models?purpose=batch"] + with running_edge(ReplayEdge(source=replay_source(root)), REPLAY_MOUNTS) as edge: + missed = call_edge(edge, "GET", "/openai/v1/models?purpose=other") + matched = call_edge(edge, "GET", "/openai/v1/models?purpose=batch") + assert missed.status_code == REPLAY_MISS_STATUS + assert matched.status_code == 200 + + def test_identical_requests_replay_their_responses_in_recorded_order(self, tmp_path: Path) -> None: + """A poll or retry loop repeats the same request and the proxy asserts + on the progression, so duplicates under one key stay FIFO.""" + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + with running_edge(ReplayEdge(source=replay_source(root)), REPLAY_MOUNTS) as edge: + first = json_object(call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")).body) + second = json_object(call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")).body) + assert first["hit"] == 1 + assert second["hit"] == 2 + + def test_exhausted_key_returns_the_miss_status(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + with running_edge(ReplayEdge(source=replay_source(root)), REPLAY_MOUNTS) as edge: + call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + exhausted = call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + assert exhausted.status_code == REPLAY_MISS_STATUS + assert b"already consumed" in exhausted.body + + def test_non_json_bodies_match_by_canonical_digest_without_storing_them(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + opaque = b"custom_id one\ncustom_id two\n" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge(edge, "POST", "/openai/v1/files", body=opaque) + raw = this_tests_files(root)[0].read_text(encoding="utf-8") + interaction = Interaction.model_validate_json(raw) + assert interaction.request.body is None + assert interaction.request.file_sha256 is not None + assert interaction.request.file_bytes == len(opaque) + assert "custom_id" not in interaction.request.model_dump_json() + with running_edge(ReplayEdge(source=replay_source(root)), REPLAY_MOUNTS) as edge: + replayed = call_edge(edge, "POST", "/openai/v1/files", body=opaque) + assert replayed.status_code == 200 + + +class TestReplayLeftover: + def test_partially_consumed_recording_names_the_leftover(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + call_edge(edge, "GET", "/openai/v1/models") + source = replay_source(root) + with running_edge(ReplayEdge(source=source), REPLAY_MOUNTS) as edge: + call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + error = source.leftover_error(current_test_key()) + assert error is not None + assert "1 of 2 recorded interactions never consumed" in error + assert "e.g. get /openai/v1/models #" in error + assert "re-record with E2E_FIXTURE_MODE=record" in error + + def test_fully_consumed_recording_leaves_nothing(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + with fake_provider() as provider: + with running_edge(record_backend(root), {"openai": provider_url(provider)}) as edge: + call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + source = replay_source(root) + with running_edge(ReplayEdge(source=source), REPLAY_MOUNTS) as edge: + call_edge(edge, "POST", CHAT_PATH, body=chat_body("hi")) + assert source.leftover_error(current_test_key()) is None + + def test_test_without_recordings_has_no_leftover(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + assert isinstance(prepare_bundle(root), BundleRecorder) + assert replay_source(root).leftover_error("suite.py::test_never_recorded") is None + + def test_inert_outside_replay_mode(self, tmp_path: Path) -> None: + missing = tmp_path / "missing" + assert replay_leftover_error(mode_raw="", bundle_dir=missing, test_key="k") is None + assert replay_leftover_error(mode_raw="record", bundle_dir=missing, test_key="k") is None + + +class TestConcurrentReplay: + def test_parallel_identical_calls_serve_each_recording_exactly_once(self, tmp_path: Path) -> None: + """The edge server handles requests on concurrent threads and a burst + of parallel identical calls consumes one shared pool: no response + duplicated, none forgotten, nothing left over at teardown.""" + root = tmp_path / "bundle" + recorder = prepare_bundle(root) + assert isinstance(recorder, BundleRecorder) + for ordinal in range(32): + recorder.record( + test_key=current_test_key(), + request=RecordedRequest(method="post", path=CHAT_PATH, headers={}, body={"n": "same"}), + response=RecordedHttpResponse( + status_code=200, + headers={"content-type": "application/json"}, + body_b64=base64.b64encode(json.dumps({"value": f"v{ordinal:02d}"}).encode()).decode(), + ), + ) + source = replay_source(root) + body = json.dumps({"n": "same"}).encode() + barrier = threading.Barrier(8) + with running_edge(ReplayEdge(source=source), REPLAY_MOUNTS) as edge: + + def consume(_: int) -> tuple[str, ...]: + barrier.wait() + return tuple( + str(json_object(call_edge(edge, "POST", CHAT_PATH, body=body).body)["value"]) + for _call in range(4) + ) + + with ThreadPoolExecutor(max_workers=8) as executor: + served = sorted(value for values in executor.map(consume, range(8)) for value in values) + assert served == [f"v{ordinal:02d}" for ordinal in range(32)] + assert source.leftover_error(current_test_key()) is None + + +class TestHandleEdgeRequestPure: + def test_unknown_mount_404s_naming_the_known_mounts(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + assert isinstance(prepare_bundle(root), BundleRecorder) + reply = handle_edge_request( + ReplayEdge(source=replay_source(root)), + {"openai": "https://api.openai.com", "anthropic": "https://api.anthropic.com"}, + "POST", + "/bedrock/model/invoke", + {}, + b"{}", + timeout=1.0, + ) + assert reply.status_code == 404 + assert b"unknown provider mount 'bedrock'" in reply.body + assert b"anthropic, openai" in reply.body + + def test_replay_serves_a_directly_recorded_interaction(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + recorder = prepare_bundle(root) + assert isinstance(recorder, BundleRecorder) + recorder.record( + test_key=current_test_key(), + request=RecordedRequest(method="post", path=CHAT_PATH, headers={}, body={"prompt": "x"}), + response=RecordedHttpResponse( + status_code=201, headers={"x-upstream": "fake"}, body_b64=base64.b64encode(b"ok").decode() + ), + ) + reply = handle_edge_request( + ReplayEdge(source=replay_source(root)), + {"openai": "https://api.openai.com"}, + "POST", + CHAT_PATH, + {"authorization": "Bearer sk-anything"}, + json.dumps({"prompt": "x"}).encode(), + timeout=1.0, + ) + assert reply.status_code == 201 + assert reply.body == b"ok" + assert reply.headers == {"x-upstream": "fake"} + + +class TestApiBaseSeam: + def test_live_mode_returns_none(self, tmp_path: Path) -> None: + for mode_raw in ("live", ""): + assert ( + provider_edge_api_base( + "openai", + mode_raw=mode_raw, + bundle_dir=tmp_path / "bundle", + bind_host="127.0.0.1", + advertise_host="127.0.0.1", + ) + is None + ) + + def test_invalid_mode_raises_naming_the_value(self, tmp_path: Path) -> None: + with pytest.raises(ValueError, match="cached"): + provider_edge_api_base( + "openai", + mode_raw="cached", + bundle_dir=tmp_path / "bundle", + bind_host="127.0.0.1", + advertise_host="127.0.0.1", + ) + + def test_unknown_mount_raises_naming_the_known_mounts(self, tmp_path: Path) -> None: + with pytest.raises(ValueError, match="unknown provider mount 'bedrock'"): + provider_edge_api_base( + "bedrock", + mode_raw="record", + bundle_dir=tmp_path / "bundle", + bind_host="127.0.0.1", + advertise_host="127.0.0.1", + ) + + def test_record_mode_boots_one_shared_edge_and_prepares_the_bundle(self, tmp_path: Path) -> None: + root = tmp_path / "bundle" + first = provider_edge_api_base( + "openai", mode_raw="record", bundle_dir=root, bind_host="127.0.0.1", advertise_host="127.0.0.1" + ) + second = provider_edge_api_base( + "anthropic", mode_raw="record", bundle_dir=root, bind_host="127.0.0.1", advertise_host="127.0.0.1" + ) + assert first is not None and second is not None + assert first.endswith("/openai") + assert second.endswith("/anthropic") + assert first.rsplit("/", 1)[0] == second.rsplit("/", 1)[0] + assert (root / "manifest.json").is_file() diff --git a/tests/enterprise/litellm_enterprise/proxy/auth/test_route_checks.py b/tests/enterprise/litellm_enterprise/proxy/auth/test_route_checks.py index c147c7aae91..f90ac9abb7d 100644 --- a/tests/enterprise/litellm_enterprise/proxy/auth/test_route_checks.py +++ b/tests/enterprise/litellm_enterprise/proxy/auth/test_route_checks.py @@ -373,3 +373,63 @@ class TestEnterpriseRouteChecksErrorMessages: # Should not raise exception for premium users result = EnterpriseRouteChecks.is_management_routes_disabled() assert result is True + + +@patch("litellm.proxy.proxy_server.premium_user", True) +class TestEnterpriseRouteChecksAgentManagement: + """Regression tests for LIT-2069: the Admin UI Agents tab could not create an + external agent on nodes with DISABLE_LLM_API_ENDPOINTS set, because agent + registry CRUD (/v1/agents*) was classified as an LLM API route. It is now a + management route, so DISABLE_ADMIN_ENDPOINTS gates it instead. Uses the real + is_llm_api_route / is_management_route classifiers (not mocks).""" + + @pytest.mark.parametrize( + "route", + [ + "/v1/agents", + "/v1/agents/abc-123", + "/v1/agents/make_public", + "/v1/agents/abc-123/make_public", + ], + ) + def test_agent_management_allowed_when_llm_api_disabled(self, route): + with patch.dict(os.environ, {"DISABLE_LLM_API_ENDPOINTS": "true"}, clear=False): + os.environ.pop("DISABLE_ADMIN_ENDPOINTS", None) + # Should not raise - agent CRUD is a management route, not llm_api. + EnterpriseRouteChecks.should_call_route(route) + + @pytest.mark.parametrize( + "route", + [ + "/v1/agents", + "/v1/agents/abc-123", + ], + ) + def test_agent_management_blocked_when_admin_disabled(self, route): + with patch.dict(os.environ, {"DISABLE_ADMIN_ENDPOINTS": "true"}, clear=False): + os.environ.pop("DISABLE_LLM_API_ENDPOINTS", None) + with pytest.raises(HTTPException) as exc_info: + EnterpriseRouteChecks.should_call_route(route) + + assert exc_info.value.status_code == 403 + assert "Management routes are disabled for this instance." in str( + exc_info.value.detail + ) + + @pytest.mark.parametrize( + "route", + [ + "/a2a/abc-123/message/send", + "/a2a/abc-123/message/stream", + ], + ) + def test_agent_inference_still_blocked_when_llm_api_disabled(self, route): + with patch.dict(os.environ, {"DISABLE_LLM_API_ENDPOINTS": "true"}, clear=False): + os.environ.pop("DISABLE_ADMIN_ENDPOINTS", None) + with pytest.raises(HTTPException) as exc_info: + EnterpriseRouteChecks.should_call_route(route) + + assert exc_info.value.status_code == 403 + assert "LLM API routes are disabled for this instance." in str( + exc_info.value.detail + ) diff --git a/tests/litellm/llms/deepseek/__init__.py b/tests/litellm/llms/deepseek/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/litellm/llms/deepseek/chat/__init__.py b/tests/litellm/llms/deepseek/chat/__init__.py deleted file mode 100644 index e69de29bb2d..00000000000 diff --git a/tests/litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py b/tests/litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py deleted file mode 100644 index 66d7e0bcbf9..00000000000 --- a/tests/litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py +++ /dev/null @@ -1,189 +0,0 @@ -""" -Unit tests for DeepSeek chat transformation. - -Tests the thinking and reasoning_effort parameter handling for DeepSeek models. -""" - -import pytest -from litellm.llms.deepseek.chat.transformation import DeepSeekChatConfig - - -class TestDeepSeekThinkingParams: - """Test thinking and reasoning_effort parameter handling for DeepSeek.""" - - def setup_method(self): - self.config = DeepSeekChatConfig() - self.model = "deepseek-reasoner" - - def test_get_supported_openai_params_includes_thinking(self): - """Test that thinking and reasoning_effort are in supported params.""" - params = self.config.get_supported_openai_params(self.model) - assert "thinking" in params - assert "reasoning_effort" in params - - def test_map_thinking_enabled(self): - """Test that thinking={"type": "enabled"} is passed through correctly.""" - non_default_params = {"thinking": {"type": "enabled"}} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - assert result["thinking"] == {"type": "enabled"} - - def test_map_thinking_with_budget_tokens_strips_budget(self): - """Test that budget_tokens is stripped from thinking param (DeepSeek doesn't support it).""" - non_default_params = {"thinking": {"type": "enabled", "budget_tokens": 2048}} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - # Should strip budget_tokens, only pass type - assert result["thinking"] == {"type": "enabled"} - assert "budget_tokens" not in result.get("thinking", {}) - - def test_map_reasoning_effort_medium(self): - """Test that reasoning_effort='medium' maps to thinking enabled.""" - non_default_params = {"reasoning_effort": "medium"} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - assert result["thinking"] == {"type": "enabled"} - - def test_map_reasoning_effort_low(self): - """Test that reasoning_effort='low' maps to thinking enabled.""" - non_default_params = {"reasoning_effort": "low"} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - assert result["thinking"] == {"type": "enabled"} - - def test_map_reasoning_effort_high(self): - """Test that reasoning_effort='high' maps to thinking enabled.""" - non_default_params = {"reasoning_effort": "high"} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - assert result["thinking"] == {"type": "enabled"} - - def test_map_reasoning_effort_none_does_not_enable_thinking(self): - """Test that reasoning_effort='none' does not enable thinking.""" - non_default_params = {"reasoning_effort": "none"} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - assert "thinking" not in result - - def test_map_reasoning_effort_null_does_not_enable_thinking(self): - """Test that reasoning_effort=None does not enable thinking.""" - non_default_params = {"reasoning_effort": None} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - assert "thinking" not in result - - def test_thinking_takes_precedence_over_reasoning_effort(self): - """Test that thinking param takes precedence when both are provided.""" - non_default_params = { - "thinking": {"type": "enabled"}, - "reasoning_effort": "high", - } - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - # thinking should be set, reasoning_effort should not override - assert result["thinking"] == {"type": "enabled"} - - def test_invalid_thinking_type_ignored(self): - """Test that invalid thinking type values are ignored.""" - non_default_params = {"thinking": {"type": "invalid"}} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - assert "thinking" not in result - - def test_thinking_none_value_ignored(self): - """Test that thinking=None is ignored.""" - non_default_params = {"thinking": None} - optional_params = {} - - result = self.config.map_openai_params( - non_default_params=non_default_params, - optional_params=optional_params, - model=self.model, - drop_params=False, - ) - - assert "thinking" not in result - - def test_drop_unsupported_tools_removes_dangling_tool_choice(self): - optional_params = { - "tools": [ - {"type": "namespace", "name": "local_shell"}, - {"type": "function", "function": {"name": "get_weather"}}, - ], - "tool_choice": { - "type": "function", - "function": {"name": "local_shell"}, - }, - "parallel_tool_calls": True, - } - - result = self.config._drop_unsupported_tools(optional_params) - - assert result["tools"] == [ - {"type": "function", "function": {"name": "get_weather"}} - ] - assert "tool_choice" not in result - assert result["parallel_tool_calls"] is True diff --git a/tests/litellm/llms/oci/chat/test_oci_chat_transformation.py b/tests/litellm/llms/oci/chat/test_oci_chat_transformation.py deleted file mode 100644 index e9b3f82d1a7..00000000000 --- a/tests/litellm/llms/oci/chat/test_oci_chat_transformation.py +++ /dev/null @@ -1,338 +0,0 @@ -""" -Tests for OCI Chat Transformation module. - -These tests verify the OCI credential handling, particularly the PEM key -normalization logic for handling different newline formats. -""" - -import os -import sys -import pytest - -sys.path.insert( - 0, os.path.abspath("../../../../..") -) # Adds the parent directory to the system path - -from litellm.llms.oci.chat.transformation import OCIChatConfig -from litellm.llms.oci.common_utils import OCIError, sign_with_manual_credentials - - -@pytest.fixture -def config(): - return OCIChatConfig() - - -class TestOCIKeyNormalization: - """Tests for OCI private key content normalization.""" - - def test_oci_key_with_escaped_newlines(self, config): - """Test that escaped newlines (\\n) are converted to actual newlines.""" - # Simulate PEM content with escaped newlines (as would come from JSON/UI input) - escaped_pem = "-----BEGIN RSA PRIVATE KEY-----\\nMIIEowIBAAKCAQEA...\\n-----END RSA PRIVATE KEY-----" - - optional_params = { - "oci_user": "ocid1.user.oc1..test", - "oci_fingerprint": "aa:bb:cc:dd", - "oci_tenancy": "ocid1.tenancy.oc1..test", - "oci_region": "us-ashburn-1", - "oci_key": escaped_pem, - } - - # We can't fully test signing without a real key, but we can verify - # the error message indicates the key was processed (not a type error) - with pytest.raises(Exception) as exc_info: - sign_with_manual_credentials( - headers={}, - optional_params=optional_params, - request_data={"test": "data"}, - api_base="https://test.oci.oraclecloud.com/api", - ) - - # The error should be about key format/loading, not about type - # This confirms the string was processed and newlines were normalized - error_message = str(exc_info.value) - assert "must be a string" not in error_message.lower() - - def test_oci_key_with_crlf_newlines(self, config): - """Test that Windows-style CRLF newlines are normalized to LF.""" - # Simulate PEM content with CRLF newlines - crlf_pem = "-----BEGIN RSA PRIVATE KEY-----\r\nMIIEowIBAAKCAQEA...\r\n-----END RSA PRIVATE KEY-----" - - optional_params = { - "oci_user": "ocid1.user.oc1..test", - "oci_fingerprint": "aa:bb:cc:dd", - "oci_tenancy": "ocid1.tenancy.oc1..test", - "oci_region": "us-ashburn-1", - "oci_key": crlf_pem, - } - - with pytest.raises(Exception) as exc_info: - sign_with_manual_credentials( - headers={}, - optional_params=optional_params, - request_data={"test": "data"}, - api_base="https://test.oci.oraclecloud.com/api", - ) - - error_message = str(exc_info.value) - assert "must be a string" not in error_message.lower() - - def test_oci_key_rejects_non_string_type(self, config): - """Test that non-string oci_key values raise OCIError.""" - optional_params = { - "oci_user": "ocid1.user.oc1..test", - "oci_fingerprint": "aa:bb:cc:dd", - "oci_tenancy": "ocid1.tenancy.oc1..test", - "oci_region": "us-ashburn-1", - "oci_key": {"invalid": "dict"}, # Wrong type - } - - with pytest.raises(OCIError) as exc_info: - sign_with_manual_credentials( - headers={}, - optional_params=optional_params, - request_data={"test": "data"}, - api_base="https://test.oci.oraclecloud.com/api", - ) - - assert exc_info.value.status_code == 400 - assert "must be a string" in str(exc_info.value.message) - assert "dict" in str(exc_info.value.message) - - def test_oci_key_rejects_list_type(self, config): - """Test that list oci_key values raise OCIError.""" - optional_params = { - "oci_user": "ocid1.user.oc1..test", - "oci_fingerprint": "aa:bb:cc:dd", - "oci_tenancy": "ocid1.tenancy.oc1..test", - "oci_region": "us-ashburn-1", - "oci_key": ["invalid", "list"], # Wrong type - } - - with pytest.raises(OCIError) as exc_info: - sign_with_manual_credentials( - headers={}, - optional_params=optional_params, - request_data={"test": "data"}, - api_base="https://test.oci.oraclecloud.com/api", - ) - - assert exc_info.value.status_code == 400 - assert "must be a string" in str(exc_info.value.message) - assert "list" in str(exc_info.value.message) - - -class TestOCIValidateEnvironment: - """Tests for OCI environment validation.""" - - def test_missing_required_credentials_raises_error(self, config): - """Test that missing required credentials raise an error.""" - with pytest.raises(Exception) as exc_info: - config.validate_environment( - headers={}, - model="oci/xai.grok-3", - messages=[{"role": "user", "content": "Hello"}], - optional_params={}, # No credentials provided - litellm_params={}, - api_key=None, - api_base=None, - ) - - error_message = str(exc_info.value) - assert "oci_user" in error_message - assert "oci_fingerprint" in error_message - assert "oci_tenancy" in error_message - - def test_validate_environment_with_all_credentials(self, config): - """Test that validation passes with all required credentials.""" - headers = config.validate_environment( - headers={}, - model="oci/xai.grok-3", - messages=[{"role": "user", "content": "Hello"}], - optional_params={ - "oci_user": "ocid1.user.oc1..test", - "oci_fingerprint": "aa:bb:cc:dd", - "oci_tenancy": "ocid1.tenancy.oc1..test", - "oci_region": "us-ashburn-1", - "oci_compartment_id": "ocid1.compartment.oc1..test", - "oci_key": "-----BEGIN RSA PRIVATE KEY-----\ntest\n-----END RSA PRIVATE KEY-----", - }, - litellm_params={}, - api_key=None, - api_base=None, - ) - - assert headers["content-type"] == "application/json" - assert "user-agent" in headers - - -class TestOCIGetCompleteUrl: - """Tests for OCI URL generation.""" - - def test_get_complete_url_default_region(self, config): - """Test URL generation with default region.""" - url = config.get_complete_url( - api_base=None, - api_key=None, - model="oci/xai.grok-3", - optional_params={}, - litellm_params={}, - stream=False, - ) - - assert "us-ashburn-1" in url - assert "inference.generativeai" in url - assert "/20231130/actions/chat" in url - - def test_get_complete_url_custom_region(self, config): - """Test URL generation with custom region.""" - url = config.get_complete_url( - api_base=None, - api_key=None, - model="oci/xai.grok-3", - optional_params={"oci_region": "eu-frankfurt-1"}, - litellm_params={}, - stream=False, - ) - - assert "eu-frankfurt-1" in url - assert "inference.generativeai" in url - - -class TestOCIImageUrlTransformation: - """Tests for OCI image_url format handling in multimodal messages. - - Fixes: https://github.com/BerriAI/litellm/issues/18270 - Fixes: https://github.com/BerriAI/litellm/issues/19589 - """ - - def test_image_url_as_string(self): - """Test that image_url as a plain string works.""" - from litellm.llms.oci.chat.transformation import ( - adapt_messages_to_generic_oci_standard, - ) - - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is in this image?"}, - {"type": "image_url", "image_url": "https://example.com/image.png"}, - ], - } - ] - - result = adapt_messages_to_generic_oci_standard(messages) - - assert len(result) == 1 - assert result[0].role == "USER" - assert len(result[0].content) == 2 - # imageUrl is now an OCIImageUrl object with a 'url' property - assert result[0].content[1].imageUrl.url == "https://example.com/image.png" - - def test_image_url_as_openai_object(self): - """Test that image_url as OpenAI-style object {"url": "..."} works.""" - from litellm.llms.oci.chat.transformation import ( - adapt_messages_to_generic_oci_standard, - ) - - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is in this image?"}, - { - "type": "image_url", - "image_url": {"url": "https://example.com/image.png"}, - }, - ], - } - ] - - result = adapt_messages_to_generic_oci_standard(messages) - - assert len(result) == 1 - assert result[0].role == "USER" - assert len(result[0].content) == 2 - # imageUrl is now an OCIImageUrl object with a 'url' property - assert result[0].content[1].imageUrl.url == "https://example.com/image.png" - - def test_image_url_serializes_as_object(self): - """Test that imageUrl serializes as {"url": "..."} for OCI API. - - Fixes: https://github.com/BerriAI/litellm/issues/19589 - OCI expects imageUrl to be an object with a 'url' property, not a plain string. - """ - from litellm.llms.oci.chat.transformation import ( - adapt_messages_to_generic_oci_standard, - ) - - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "Describe this image."}, - { - "type": "image_url", - "image_url": {"url": "data:image/png;base64,ABC123"}, - }, - ], - } - ] - - result = adapt_messages_to_generic_oci_standard(messages) - image_part = result[0].content[1] - - # Serialize as OCI would receive it (with exclude_none=True) - serialized = image_part.model_dump(exclude_none=True) - - # Verify the structure matches OCI's expected format - assert serialized == { - "type": "IMAGE", - "imageUrl": {"url": "data:image/png;base64,ABC123"}, - } - - def test_image_url_invalid_type_raises_error(self): - """Test that invalid image_url type raises an error.""" - from litellm.llms.oci.chat.transformation import ( - adapt_messages_to_generic_oci_standard, - ) - - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is in this image?"}, - {"type": "image_url", "image_url": 12345}, # Invalid type - ], - } - ] - - with pytest.raises(Exception) as exc_info: - adapt_messages_to_generic_oci_standard(messages) - - assert "image_url" in str(exc_info.value) - - def test_image_url_object_missing_url_raises_error(self): - """Test that object without 'url' property raises an error.""" - from litellm.llms.oci.chat.transformation import ( - adapt_messages_to_generic_oci_standard, - ) - - messages = [ - { - "role": "user", - "content": [ - {"type": "text", "text": "What is in this image?"}, - { - "type": "image_url", - "image_url": {"detail": "high"}, - }, # Missing 'url' - ], - } - ] - - with pytest.raises(Exception) as exc_info: - adapt_messages_to_generic_oci_standard(messages) - - assert "image_url" in str(exc_info.value) diff --git a/tests/litellm/proxy/management_endpoints/test_common_utils.py b/tests/litellm/proxy/management_endpoints/test_common_utils.py deleted file mode 100644 index f857db770d0..00000000000 --- a/tests/litellm/proxy/management_endpoints/test_common_utils.py +++ /dev/null @@ -1,159 +0,0 @@ -""" -Tests for litellm/proxy/management_endpoints/common_utils.py - -Specifically tests that _update_metadata_fields does not trigger premium -user checks when premium fields are present but empty. - -Related: https://github.com/BerriAI/litellm/issues/20534 -""" - -from unittest.mock import patch - -import pytest - -from litellm.proxy.management_endpoints.common_utils import ( - _has_non_empty_value, - _update_metadata_fields, -) - - -class TestHasNonEmptyValue: - """Tests for the _has_non_empty_value helper.""" - - def test_none_is_empty(self): - assert _has_non_empty_value(None) is False - - def test_empty_list_is_empty(self): - assert _has_non_empty_value([]) is False - - def test_empty_string_is_empty(self): - assert _has_non_empty_value("") is False - - def test_blank_string_is_empty(self): - assert _has_non_empty_value(" ") is False - - def test_non_empty_list_has_value(self): - assert _has_non_empty_value(["policy-a"]) is True - - def test_non_empty_string_has_value(self): - assert _has_non_empty_value("30d") is True - - def test_dict_has_value(self): - assert _has_non_empty_value({"key": "val"}) is True - - def test_empty_dict_has_value(self): - # empty dict is not None/list/str, so it counts as non-empty - assert _has_non_empty_value({}) is True - - -class TestUpdateMetadataFieldsPremiumCheck: - """ - Tests that _update_metadata_fields skips premium user checks for empty - values but still enforces them for real values. - - Issue: The UI sends the full form on every team update, including premium - fields like `policies: []`. The backend was treating these empty values - as premium feature usage and returning 403. - """ - - @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", - side_effect=Exception("Should not be called"), - ) - def test_empty_policies_skips_premium_check(self, mock_check): - """policies: [] should NOT trigger premium user check.""" - updated_kv = { - "team_id": "team-123", - "team_alias": "my-team", - "policies": [], - } - _update_metadata_fields(updated_kv) - mock_check.assert_not_called() - - @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", - side_effect=Exception("Should not be called"), - ) - def test_empty_guardrails_skips_premium_check(self, mock_check): - """guardrails: [] should NOT trigger premium user check.""" - updated_kv = { - "team_id": "team-123", - "guardrails": [], - } - _update_metadata_fields(updated_kv) - mock_check.assert_not_called() - - @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", - side_effect=Exception("Should not be called"), - ) - def test_empty_string_team_member_key_duration_skips_premium_check( - self, mock_check - ): - """team_member_key_duration: '' should NOT trigger premium user check.""" - updated_kv = { - "team_id": "team-123", - "team_member_key_duration": "", - } - _update_metadata_fields(updated_kv) - mock_check.assert_not_called() - - @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", - side_effect=Exception("Should not be called"), - ) - def test_full_ui_payload_with_empty_premium_fields_skips_premium_check( - self, mock_check - ): - """A realistic UI payload with all empty premium fields should not 403.""" - updated_kv = { - "team_id": "team-123", - "team_alias": "renamed-team", - "models": ["gpt-4o"], - "max_budget": 200, - "policies": [], - "guardrails": [], - "logging": [], - "team_member_key_duration": "", - "prompts": [], - } - _update_metadata_fields(updated_kv) - mock_check.assert_not_called() - - @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", - ) - def test_non_empty_policies_triggers_premium_check(self, mock_check): - """policies: ['real-policy'] SHOULD trigger premium user check.""" - updated_kv = { - "team_id": "team-123", - "policies": ["real-policy"], - } - _update_metadata_fields(updated_kv) - mock_check.assert_called() - - @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", - ) - def test_non_empty_guardrails_triggers_premium_check(self, mock_check): - """guardrails: ['my-guardrail'] SHOULD trigger premium user check.""" - updated_kv = { - "team_id": "team-123", - "guardrails": ["my-guardrail"], - } - _update_metadata_fields(updated_kv) - mock_check.assert_called() - - @patch( - "litellm.proxy.management_endpoints.common_utils._premium_user_check", - ) - def test_non_empty_team_member_key_duration_triggers_premium_check( - self, mock_check - ): - """team_member_key_duration: '30d' SHOULD trigger premium user check.""" - updated_kv = { - "team_id": "team-123", - "team_member_key_duration": "30d", - } - _update_metadata_fields(updated_kv) - mock_check.assert_called() diff --git a/tests/litellm_utils_tests/test_health_check.py b/tests/litellm_utils_tests/test_health_check.py index de6f7c38fed..654fde90f26 100644 --- a/tests/litellm_utils_tests/test_health_check.py +++ b/tests/litellm_utils_tests/test_health_check.py @@ -132,7 +132,7 @@ async def test_azure_img_gen_health_check(): retry_delay *= 2 # Exponential backoff # Should not reach here, but just in case - assert False, "Health check failed after all retries" + pytest.fail("Health check failed after all retries") @pytest.mark.skip(reason="AWS Suspended Account") diff --git a/tests/llm_translation/test_bedrock_completion.py b/tests/llm_translation/test_bedrock_completion.py index 8dcae7cc997..3303fafafb0 100644 --- a/tests/llm_translation/test_bedrock_completion.py +++ b/tests/llm_translation/test_bedrock_completion.py @@ -2442,9 +2442,7 @@ class TestBedrockEmbedding(BaseLLMEmbeddingTest): transformed_request = ( AmazonTitanMultimodalEmbeddingG1Config()._transform_request(**args) ) - transformed_request[ - "inputImage" - ] == "iVBORw0KGgoAAAANSUhEUgAAAGQAAABkBAMAAACCzIhnAAAAG1BMVEURAAD///+ln5/h39/Dv79qX18uHx+If39MPz9oMSdmAAAACXBIWXMAAA7EAAAOxAGVKw4bAAABB0lEQVRYhe2SzWrEIBCAh2A0jxEs4j6GLDS9hqWmV5Flt0cJS+lRwv742DXpEjY1kOZW6HwHFZnPmVEBEARBEARB/jd0KYA/bcUYbPrRLh6amXHJ/K+ypMoyUaGthILzw0l+xI0jsO7ZcmCcm4ILd+QuVYgpHOmDmz6jBeJImdcUCmeBqQpuqRIbVmQsLCrAalrGpfoEqEogqbLTWuXCPCo+Ki1XGqgQ+jVVuhB8bOaHkvmYuzm/b0KYLWwoK58oFqi6XfxQ4Uz7d6WeKpna6ytUs5e8betMcqAv5YPC5EZB2Lm9FIn0/VP6R58+/GEY1X1egVoZ/3bt/EqF6malgSAIgiDIH+QL41409QMY0LMAAAAASUVORK5CYII=" + assert transformed_request["inputImage"] == "iVBORw0KGgoAAAANSUhEUgAAAGQAAABkBAMAAACCzIhnAAAAG1BMVEURAAD///+ln5/h39/Dv79qX18uHx+If39MPz9oMSdmAAAACXBIWXMAAA7EAAAOxAGVKw4bAAABB0lEQVRYhe2SzWrEIBCAh2A0jxEs4j6GLDS9hqWmV5Flt0cJS+lRwv742DXpEjY1kOZW6HwHFZnPmVEBEARBEARB/jd0KYA/bcUYbPrRLh6amXHJ/K+ypMoyUaGthILzw0l+xI0jsO7ZcmCcm4ILd+QuVYgpHOmDmz6jBeJImdcUCmeBqQpuqRIbVmQsLCrAalrGpfoEqEogqbLTWuXCPCo+Ki1XGqgQ+jVVuhB8bOaHkvmYuzm/b0KYLWwoK58oFqi6XfxQ4Uz7d6WeKpna6ytUs5e8betMcqAv5YPC5EZB2Lm9FIn0/VP6R58+/GEY1X1egVoZ/3bt/EqF6malgSAIgiDIH+QL41409QMY0LMAAAAASUVORK5CYII=" @pytest.mark.asyncio diff --git a/tests/llm_translation/test_bedrock_embedding.py b/tests/llm_translation/test_bedrock_embedding.py index 2bc4192833b..e343b8856a7 100644 --- a/tests/llm_translation/test_bedrock_embedding.py +++ b/tests/llm_translation/test_bedrock_embedding.py @@ -447,9 +447,7 @@ def test_bedrock_embedding_region_bug_reproduction(): print( "❌ BUG REPRODUCED: Using wrong region from env var instead of explicit parameter" ) - assert ( - False - ), f"Bug reproduced: URL contains ap-northeast-1 instead of us-east-1. URL: {url}" + pytest.fail(f"Bug reproduced: URL contains ap-northeast-1 instead of us-east-1. URL: {url}") else: print( "✓ Bug NOT reproduced: Using correct region from explicit parameter" diff --git a/tests/llm_translation/test_containers_api.py b/tests/llm_translation/test_containers_api.py index 2ae93a3a406..4e26a883fcb 100644 --- a/tests/llm_translation/test_containers_api.py +++ b/tests/llm_translation/test_containers_api.py @@ -70,7 +70,7 @@ def test_container_files_api(): custom_llm_provider="openai", api_key=api_key, ) - assert False, "Should have raised error for non-existent file" + pytest.fail("Should have raised error for non-existent file") except Exception as e: assert "not found" in str(e).lower() or "invalid" in str(e).lower() print(f" Got expected error ✓") @@ -84,7 +84,7 @@ def test_container_files_api(): custom_llm_provider="openai", api_key=api_key, ) - assert False, "Should have raised error for non-existent file content" + pytest.fail("Should have raised error for non-existent file content") except Exception as e: print(f" Got expected error ✓") @@ -97,7 +97,7 @@ def test_container_files_api(): custom_llm_provider="openai", api_key=api_key, ) - assert False, "Should have raised error for non-existent file" + pytest.fail("Should have raised error for non-existent file") except Exception as e: # Delete returns 400 for non-existent files print(f" Got expected error ✓") diff --git a/tests/llm_translation/test_gemini.py b/tests/llm_translation/test_gemini.py index 0a5aebdf91b..9326e291b7f 100644 --- a/tests/llm_translation/test_gemini.py +++ b/tests/llm_translation/test_gemini.py @@ -1324,7 +1324,7 @@ def test_gemini_exception_message_format(): extra_kwargs={}, ) # Should not reach here - exception should be raised - assert False, "Expected BadRequestError to be raised" + pytest.fail("Expected BadRequestError to be raised") except BadRequestError as e: # The test should FAIL initially (before fix) because it will show VertexAIException # After the fix, it should show GeminiException @@ -1401,9 +1401,7 @@ def l(status_code, expected_exception): completion_kwargs={}, extra_kwargs={}, ) - assert ( - False - ), f"Expected {expected_exception} to be raised for status {status_code}" + pytest.fail(f"Expected {expected_exception} to be raised for status {status_code}") except Exception as e: # Verify the correct exception type is raised exception_classes = { diff --git a/tests/local_testing/test_completion.py b/tests/local_testing/test_completion.py index 6f58bb2eb35..01fd35cb42d 100644 --- a/tests/local_testing/test_completion.py +++ b/tests/local_testing/test_completion.py @@ -842,6 +842,8 @@ def test_completion_mistral_api_modified_input(): @pytest.mark.skip(reason="this test is flaky") def test_completion_gpt4_vision(): + import openai + try: litellm.set_verbose = True response = completion( @@ -1820,6 +1822,8 @@ def test_completion_openai_litellm_key(): @pytest.mark.skip(reason="Unresponsive endpoint.[TODO] Rehost this somewhere else") def test_completion_ollama_hosted(): + import openai + try: litellm.request_timeout = 20 # give ollama 20 seconds to response litellm.set_verbose = True @@ -2057,17 +2061,12 @@ def test_completion_openrouter_reasoning_effort(): def test_completion_hf_model_no_provider(): - try: - response = completion( + with pytest.raises(litellm.BadRequestError, match="LLM Provider NOT provided"): + completion( model="WizardLM/WizardLM-70B-V1.0", messages=messages, max_tokens=5, ) - # Add any assertions here to check the response - print(response) - pytest.fail(f"Error occurred: {e}") - except Exception as e: - pass # test_completion_hf_model_no_provider() @@ -2546,7 +2545,7 @@ def test_completion_replicate_vicuna(): response_str = response["choices"][0]["message"]["content"] print("RESPONSE STRING\n", response_str) if type(response_str) != str: - pytest.fail(f"Error occurred: {e}") + pytest.fail(f"Expected a string response, got {type(response_str)}: {response_str}") except Exception as e: pytest.fail(f"Error occurred: {e}") diff --git a/tests/local_testing/test_completion_cost.py b/tests/local_testing/test_completion_cost.py index 551b3f064bb..e34d5c349c5 100644 --- a/tests/local_testing/test_completion_cost.py +++ b/tests/local_testing/test_completion_cost.py @@ -1097,7 +1097,7 @@ def test_completion_cost_databricks(model): litellm._turn_on_debug() os.environ["LITELLM_LOCAL_MODEL_COST_MAP"] = "True" litellm.model_cost = litellm.get_model_cost_map(url="") - model, messages = model, [{"role": "user", "content": "What is 2+2?"}] + messages = [{"role": "user", "content": "What is 2+2?"}] resp = litellm.completion(model=model, messages=messages) # works fine @@ -1479,7 +1479,6 @@ def test_completion_cost_azure_ai_rerank(model): }, ) print("response", response) - model = model cost = completion_cost( model=model, completion_response=response, call_type="arerank" ) @@ -2874,7 +2873,7 @@ def test_json_valid_model_cost_map(): json_str = json.dumps(model_cost) json.loads(json_str) except json.JSONDecodeError as e: - assert False, f"Invalid JSON format: {str(e)}" + pytest.fail(f"Invalid JSON format: {str(e)}") def test_batch_cost_calculator(): diff --git a/tests/local_testing/test_custom_llm.py b/tests/local_testing/test_custom_llm.py index ea15c3db9d0..64a6c8b2587 100644 --- a/tests/local_testing/test_custom_llm.py +++ b/tests/local_testing/test_custom_llm.py @@ -489,9 +489,9 @@ async def test_image_generation_async_additional_params(): mock_client.assert_awaited_once() - mock_client.call_args.kwargs["api_key"] == "my-api-key" - mock_client.call_args.kwargs["api_base"] == "my-api-base" - mock_client.call_args.kwargs["optional_params"] == { + assert mock_client.call_args.kwargs["api_key"] == "my-api-key" + assert mock_client.call_args.kwargs["api_base"] == "my-api-base" + assert mock_client.call_args.kwargs["optional_params"] == { "my_custom_param": "my-custom-param" } diff --git a/tests/local_testing/test_custom_logger.py b/tests/local_testing/test_custom_logger.py index 6af2ff7e964..02a9eaaa9e6 100644 --- a/tests/local_testing/test_custom_logger.py +++ b/tests/local_testing/test_custom_logger.py @@ -279,7 +279,6 @@ def test_azure_completion_stream(): @pytest.mark.asyncio async def test_async_custom_handler_completion(): try: - litellm._turn_on_debug customHandler_success = MyCustomHandler() customHandler_failure = MyCustomHandler() # success diff --git a/tests/local_testing/test_exceptions.py b/tests/local_testing/test_exceptions.py index e02d9e21171..8c1df52e28e 100644 --- a/tests/local_testing/test_exceptions.py +++ b/tests/local_testing/test_exceptions.py @@ -573,7 +573,7 @@ def test_content_policy_violation_error_streaming(): num_finish_reason += 1 print("finish_reason", chunk["choices"][0].get("finish_reason")) - pytest.fail(f"Expected to return 400 error In streaming{e}") + pytest.fail("Expected a content-policy error in streaming, got a clean stream") except Exception as e: pass diff --git a/tests/local_testing/test_fake_openai_endpoint.py b/tests/local_testing/test_fake_openai_endpoint.py index 79d8b4f97e3..d5236d3de1b 100644 --- a/tests/local_testing/test_fake_openai_endpoint.py +++ b/tests/local_testing/test_fake_openai_endpoint.py @@ -13,9 +13,11 @@ from __future__ import annotations import re from pathlib import Path +from typing import Final import httpx import pytest +from openai.types import ModerationCreateResponse from tests.fake_openai_endpoint import ( _LOCAL_DEFAULT, @@ -56,6 +58,20 @@ def test_chat_completion_shape(): assert body["usage"]["total_tokens"] == 40 +def test_moderations_route_parses_as_an_openai_response(): + base: Final = ensure_fake_openai_endpoint() + response: Final = httpx.post( + f"{base}/v1/moderations", + json={"input": ["I want to harm someone", "hello"], "model": "omni-moderation-latest"}, + timeout=10, + ) + assert response.status_code == 200 + parsed: Final = ModerationCreateResponse.model_validate(response.json()) + assert parsed.model == "omni-moderation-latest" + assert len(parsed.results) == 2 + assert parsed.results[0].categories.violence is False + + def test_triton_embeddings_route(): base = ensure_fake_openai_endpoint() response = httpx.post(f"{base}/triton/embeddings", json={"inputs": []}, timeout=10) diff --git a/tests/local_testing/test_llm_guard.py b/tests/local_testing/test_llm_guard.py index 78bbd1c0af8..86fa80ee944 100644 --- a/tests/local_testing/test_llm_guard.py +++ b/tests/local_testing/test_llm_guard.py @@ -15,6 +15,8 @@ sys.path.insert( 0, os.path.abspath("../..") ) # Adds the parent directory to the system path import pytest +from fastapi import HTTPException + import litellm from litellm_enterprise.enterprise_callbacks.llm_guard import _ENTERPRISE_LLMGuard from litellm import Router, mock_completion @@ -128,7 +130,7 @@ async def test_llm_guard_error_raising(): user_api_key_dict = UserAPIKeyAuth(api_key=_api_key) local_cache = DualCache() - try: + with pytest.raises(HTTPException) as exc_info: await llm_guard.async_moderation_hook( data={ "messages": [ @@ -141,9 +143,9 @@ async def test_llm_guard_error_raising(): user_api_key_dict=user_api_key_dict, call_type="completion", ) - pytest.fail(f"Should have failed - {str(e)}") - except Exception as e: - pass + + assert exc_info.value.status_code == 400 + assert exc_info.value.detail == {"error": "Violated content safety policy"} def test_llm_guard_key_specific_mode(): diff --git a/tests/local_testing/test_router_pattern_matching.py b/tests/local_testing/test_router_pattern_matching.py index d09790d43b1..d02582a2a99 100644 --- a/tests/local_testing/test_router_pattern_matching.py +++ b/tests/local_testing/test_router_pattern_matching.py @@ -5,6 +5,7 @@ Pattern matching router is used to match patterns like openai/*, vertex_ai/*, an """ import sys, os, time +import json import traceback, asyncio import pytest @@ -233,11 +234,11 @@ def test_router_pattern_match_e2e(): api_key="test", ) mock_post.assert_called_once() - print(mock_post.call_args.kwargs["data"]) - mock_post.call_args.kwargs["data"] == { - "model": "gpt-4o", - "messages": [{"role": "user", "content": "Hello, how are you?"}], - } + request_body = json.loads(mock_post.call_args.kwargs["data"]) + assert request_body["model"] == "my-custom-model" + assert request_body["messages"] == [ + {"role": "user", "content": [{"type": "text", "text": "Hello, how are you?"}]} + ] def test_pattern_matching_router_with_default_wildcard(): diff --git a/tests/local_testing/test_secret_detect_hook.py b/tests/local_testing/test_secret_detect_hook.py index 57b55bd2689..ad2e248da1b 100644 --- a/tests/local_testing/test_secret_detect_hook.py +++ b/tests/local_testing/test_secret_detect_hook.py @@ -137,7 +137,7 @@ async def test_basic_secret_detection_text_completion(): call_type="completion", ) - test_data == { + assert test_data == { "prompt": "Hey, how's it going, API_KEY = '[REDACTED]', my OPENAI_API_KEY = '[REDACTED]' and i want to know what is the weather", "model": "gpt-3.5-turbo", } diff --git a/tests/logging_callback_tests/test_gcs_pub_sub.py b/tests/logging_callback_tests/test_gcs_pub_sub.py index c37a2e3f65d..3c242a5fe1d 100644 --- a/tests/logging_callback_tests/test_gcs_pub_sub.py +++ b/tests/logging_callback_tests/test_gcs_pub_sub.py @@ -133,7 +133,7 @@ def assert_gcs_pubsub_request_matches_expected( actual_request_body, expected_request_body, ignore_keys=ignored_keys ) if differences: - assert False, f"Dictionary mismatch: {differences}" + pytest.fail(f"Dictionary mismatch: {differences}") def assert_gcs_pubsub_request_matches_expected_standard_logging_payload( diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py index 32295310005..6da8ce598a9 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py @@ -1441,9 +1441,7 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e(): print( f"ERROR: Duplicate MCP fetching detected! Called {mock_get_tools.call_count} times" ) - assert ( - False - ), f"MCP tools should be fetched exactly once, but were fetched {mock_get_tools.call_count} times" + pytest.fail(f"MCP tools should be fetched exactly once, but were fetched {mock_get_tools.call_count} times") # Additional validation: ensure no duplicate tools in any LLM call total_duplicates_found = 0 @@ -1466,9 +1464,7 @@ async def test_no_duplicate_mcp_tools_in_streaming_e2e(): ) if total_duplicates_found > 0: - assert ( - False - ), f"Found {total_duplicates_found} duplicate tools across all LLM calls" + pytest.fail(f"Found {total_duplicates_found} duplicate tools across all LLM calls") print("No duplicate MCP tools E2E test passed!") print(f"Summary:") diff --git a/tests/proxy_e2e_anthropic_messages_tests/test_claude_agent_sdk.py b/tests/proxy_e2e_anthropic_messages_tests/test_claude_agent_sdk.py index 48eb7d85ec1..c1339ce6280 100644 --- a/tests/proxy_e2e_anthropic_messages_tests/test_claude_agent_sdk.py +++ b/tests/proxy_e2e_anthropic_messages_tests/test_claude_agent_sdk.py @@ -148,65 +148,6 @@ async def test_claude_agent_sdk_streaming( f"Test failed for {model_name} ({model_description}) after {MAX_RETRIES} attempts: {last_error}" ) - # Test query - test_query = "Say 'Hello from LiteLLM!' and nothing else." - - # Track streaming - received_chunks = [] - full_response = "" - - try: - async with ClaudeSDKClient(options=options) as client: - await client.query(test_query) - - # Collect streaming response - async for msg in client.receive_response(): - # Handle different message types - if hasattr(msg, "type"): - if msg.type == "content_block_delta": - # Streaming text delta - if hasattr(msg, "delta") and hasattr(msg.delta, "text"): - chunk_text = msg.delta.text - received_chunks.append(chunk_text) - full_response += chunk_text - elif msg.type == "content_block_start": - # Start of content block - if hasattr(msg, "content_block") and hasattr( - msg.content_block, "text" - ): - chunk_text = msg.content_block.text - received_chunks.append(chunk_text) - full_response += chunk_text - - # Fallback to content handling - if hasattr(msg, "content"): - for content_block in msg.content: - if hasattr(content_block, "text"): - chunk_text = content_block.text - received_chunks.append(chunk_text) - full_response += chunk_text - - # Assertions - print(f"\n✅ Received {len(received_chunks)} chunks") - print(f"📝 Full response: {full_response[:100]}...") - - # Verify we got a response - assert len(full_response) > 0, f"No response received from {model_name}" - - # Verify streaming (should have multiple chunks for most responses) - # Note: Very short responses might come in 1 chunk, so we just verify we got content - assert len(received_chunks) > 0, f"No chunks received from {model_name}" - - # Verify response is non-empty (don't assert on specific LLM content — it's non-deterministic) - assert ( - len(full_response.strip()) > 0 - ), f"Empty response received from {model_name}" - - print(f"✅ Test passed for {model_name}") - - except Exception as e: - pytest.fail(f"Test failed for {model_name} ({model_description}): {str(e)}") - if __name__ == "__main__": # Run tests diff --git a/tests/proxy_unit_tests/test_check_batch_cost.py b/tests/proxy_unit_tests/test_check_batch_cost.py index 1dbbbfc43a0..0065dbebc59 100644 --- a/tests/proxy_unit_tests/test_check_batch_cost.py +++ b/tests/proxy_unit_tests/test_check_batch_cost.py @@ -6,11 +6,17 @@ Vertex (raw gs:// input_file_id) and Bedrock (raw s3:// input_file_id, ARN unified_object_id) batches with no managed unified id. """ +import asyncio +import json +from contextlib import contextmanager from unittest.mock import AsyncMock, MagicMock, patch import pytest +from fastapi import HTTPException _IS_B64 = "litellm.proxy.openai_files_endpoints.common_utils._is_base64_encoded_unified_file_id" +_CLAIM_UNIFIED_BATCH_ID = "dW5pZmllZF9iYXRjaF9pZA==" +_CLAIM_OUTPUT_FILE_ID = "file-output-123" def _unmanaged_vertex_file_object( @@ -95,7 +101,7 @@ class TestCheckBatchCost: ): """_cleanup_stale_managed_objects scopes its update to file_purpose='batch' only.""" mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Return empty so the main poll loop exits immediately mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( @@ -161,7 +167,7 @@ class TestCheckBatchCost: from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( return_value=[] @@ -192,7 +198,7 @@ class TestCheckBatchCost: from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # First find_many (primary query) raises with a schema error; second (fallback) returns empty mock_prisma_client.db.litellm_managedobjecttable.find_many = AsyncMock( @@ -221,7 +227,7 @@ class TestCheckBatchCost: from litellm.constants import MAX_OBJECTS_PER_POLL_CYCLE mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) # Simulate column already known absent from a previous cycle check_batch_cost_instance._has_batch_processed_column = False @@ -254,7 +260,7 @@ class TestCheckBatchCost: from unittest.mock import patch mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -563,7 +569,7 @@ class TestCheckBatchCost: from unittest.mock import patch mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -679,7 +685,7 @@ class TestCheckBatchCost: import litellm from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger - mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0) + mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock(return_value=None) @@ -801,7 +807,7 @@ class TestCheckBatchCost: from unittest.mock import patch mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -869,7 +875,7 @@ class TestCheckBatchCost: import base64 mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -944,7 +950,7 @@ class TestCheckBatchCost: ).decode() mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -1044,7 +1050,7 @@ class TestCheckBatchCost: from unittest.mock import patch mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -1111,7 +1117,7 @@ class TestCheckBatchCost: from unittest.mock import patch mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() @@ -1168,7 +1174,7 @@ class TestCheckBatchCost: from unittest.mock import patch mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -1284,7 +1290,7 @@ class TestCheckBatchCost: from litellm.exceptions import NotFoundError mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -1355,7 +1361,7 @@ class TestCheckBatchCost: through the proxy, causing API_KEY errors when clients call GET /files/{id}/content. """ mock_prisma_client.db.litellm_managedobjecttable.update_many = AsyncMock( - return_value=0 + return_value=1 ) mock_prisma_client.db.litellm_managedobjecttable.update = AsyncMock() mock_prisma_client.db.litellm_usertable.find_unique = AsyncMock( @@ -1672,7 +1678,7 @@ class TestUnmanagedVertexRouting: prisma = instance.prisma_client prisma.db = MagicMock() prisma.db.litellm_managedobjecttable = MagicMock() - prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0) + prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) prisma.db.litellm_managedobjecttable.update = AsyncMock() prisma.db.litellm_managedobjecttable.find_many = AsyncMock( return_value=[self._job()] @@ -1902,7 +1908,7 @@ class TestUnmanagedBedrockRouting: prisma = instance.prisma_client prisma.db = MagicMock() prisma.db.litellm_managedobjecttable = MagicMock() - prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=0) + prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) prisma.db.litellm_managedobjecttable.update = AsyncMock() prisma.db.litellm_managedobjecttable.find_many = AsyncMock( return_value=[self._job()] @@ -2577,3 +2583,353 @@ class TestPollPageStarvation: await self._instance(prisma, llm_router).check_batch_cost() prisma.db.litellm_managedobjecttable.update.assert_not_awaited() + +class _FakeManagedObjectRow: + """One managed batch row the provider has finished but nothing has costed yet.""" + + def __init__(self): + self.id = "job-claim-1" + self.unified_object_id = _CLAIM_UNIFIED_BATCH_ID + self.model_object_id = "batch-456" + self.file_purpose = "batch" + self.status = "in_progress" + self.batch_processed = False + self.created_by = "user-1" + self.team_id = None + self.api_key = None + self.request_tags = None + self.created_at = 1700000000 + self.file_object = json.dumps( + {"id": "batch-456", "status": "in_progress", "input_file_id": "file-input-1", + "output_file_id": _CLAIM_OUTPUT_FILE_ID} + ) + + +class _FakeManagedObjectTable: + """A LiteLLM_ManagedObjectTable double backed by one real, mutable row. + + It honours the batch_processed and status filters, so the poller's compare-and-swap + and the managed-files deletion guard both read the same state a shared Postgres row + would give them. Staleness sweeps (the only queries scoped by created_at) never match. + """ + + def __init__(self, row: _FakeManagedObjectRow, journal: list): + self.row = row + self.journal = journal + self.update_many = AsyncMock(side_effect=self._update_many) + self.update = AsyncMock(side_effect=self._update) + self.find_many = AsyncMock(side_effect=self._find_many) + self.find_first = AsyncMock(return_value=None) + + def _matches(self, where: dict) -> bool: + for key, value in where.items(): + if key == "created_at": + return False + if key == "status": + if self.row.status in value.get("not_in", []): + return False + if "in" in value and self.row.status not in value["in"]: + return False + elif getattr(self.row, key) != value: + return False + return True + + async def _update_many(self, *, where: dict, data: dict) -> int: + if not self._matches(where): + return 0 + if "batch_processed" in where: + self.journal.append("claim" if data.get("batch_processed") else "release") + for key, value in data.items(): + setattr(self.row, key, value) + return 1 + + async def _update(self, *, where: dict, data: dict) -> None: + self.journal.append("finalize") + for key, value in data.items(): + setattr(self.row, key, value) + + async def _find_many(self, *, where: dict, take=None, order=None) -> list: + return [self.row] if self._matches(where) else [] + + +class TestMultiPodBatchCostClaim: + """LIT-4827 regression: every pod and uvicorn worker schedules its own poller against + the shared LiteLLM_ManagedObjectTable, so a completed batch must be claimed atomically + before its cost is logged. Without the claim two pods select the same row in one window + and both write an aretrieve_batch spend log for it, double counting the spend. + + The claim sits immediately before the spend-log write rather than before the results + fetch, because batch_processed is also what keeps an unbilled row selectable by later + poll cycles and what blocks deletion of the files the fetch reads.""" + + @staticmethod + def _instance(prisma, llm_router): + from litellm_enterprise.proxy.common_utils.check_batch_cost import CheckBatchCost + + proxy_logging_obj = MagicMock() + proxy_logging_obj.get_proxy_hook.return_value = None + return CheckBatchCost( + proxy_logging_obj=proxy_logging_obj, + prisma_client=prisma, + llm_router=llm_router, + ) + + @staticmethod + def _prisma(row: _FakeManagedObjectRow, journal: list): + prisma = MagicMock() + prisma.db.litellm_managedobjecttable = _FakeManagedObjectTable(row, journal) + prisma.db.litellm_managedfiletable.find_many = AsyncMock(return_value=[]) + prisma.db.litellm_managedfiletable.find_first = AsyncMock(return_value=None) + prisma.db.litellm_usertable.find_unique = AsyncMock(return_value=None) + prisma.db.litellm_verificationtoken.find_unique = AsyncMock(return_value=None) + return prisma + + @staticmethod + def _router(): + response = MagicMock() + response.status = "completed" + response.output_file_id = _CLAIM_OUTPUT_FILE_ID + response.error_file_id = None + response.created_at = 1 + response.completed_at = 2 + response.model_dump_json.return_value = '{"id":"batch-456","status":"completed"}' + + deployment = MagicMock() + deployment.litellm_params.custom_llm_provider = "openai" + deployment.litellm_params.model = "gpt-4" + deployment.model_info.model_dump.return_value = {} + + router = MagicMock() + router.aretrieve_batch = AsyncMock(return_value=response) + router.get_deployment_credentials_with_provider = MagicMock( + return_value={"api_key": "sk-test"} + ) + router.get_deployment = MagicMock(return_value=deployment) + return router + + @staticmethod + @contextmanager + def _billing_patches(journal: list, during_fetch=None, bill_error=None): + """Patch the cost path a batch runs through, journalling the results fetch and the + spend-log write. during_fetch runs while the output file is being read, which is + the window an interrupted worker or a concurrent file deletion lands in.""" + file_content = MagicMock() + file_content.content = b'{"id":"req-1"}' + + async def _afile_content(**kwargs): + journal.append("fetch") + if during_fetch is not None: + await during_fetch() + return file_content + + async def _bill(**kwargs): + journal.append("bill") + if bill_error is not None: + raise bill_error + + def _is_b64(file_id): + if file_id == _CLAIM_UNIFIED_BATCH_ID: + return "llm_model_id,model-123;llm_batch_id,batch-456;" + return False + + logging_obj = MagicMock() + logging_obj.async_success_handler = AsyncMock(side_effect=_bill) + + with ( + patch(_IS_B64, side_effect=_is_b64), + patch( + "litellm.proxy.openai_files_endpoints.common_utils.get_model_id_from_unified_batch_id", + return_value="model-123", + ), + patch( + "litellm.proxy.openai_files_endpoints.common_utils.get_batch_id_from_unified_batch_id", + return_value="batch-456", + ), + patch("litellm.files.main.afile_content", new=AsyncMock(side_effect=_afile_content)), + patch( + "litellm.batches.batch_utils._get_file_content_as_dictionary", + return_value=[{"id": "req-1"}], + ), + patch( + "litellm.batches.batch_utils.calculate_batch_cost_and_usage", + new_callable=AsyncMock, + return_value=(0.01, {"prompt_tokens": 10, "completion_tokens": 5}, ["gpt-4"]), + ), + patch( + "litellm.litellm_core_utils.get_llm_provider_logic.get_llm_provider", + return_value=("gpt-4", "openai", None, None), + ), + patch("litellm.litellm_core_utils.litellm_logging.Logging", return_value=logging_obj), + ): + yield logging_obj + + @staticmethod + def _claim_calls(prisma) -> list: + return [ + call.kwargs + for call in prisma.db.litellm_managedobjecttable.update_many.call_args_list + if "id" in call.kwargs["where"] + ] + + @staticmethod + async def _run_deletion_guard(prisma, file_id: str) -> None: + """Run the real managed-files deletion guard against the row the poller is costing.""" + from litellm_enterprise.proxy.hooks.managed_files import _PROXY_LiteLLMManagedFiles + + cache = MagicMock() + cache.async_get_cache = AsyncMock(return_value=None) + cache.async_set_cache = AsyncMock() + guard = _PROXY_LiteLLMManagedFiles(internal_usage_cache=cache, prisma_client=prisma) + + scheduler = MagicMock() + scheduler.get_job.return_value = MagicMock() + with patch("litellm.proxy.proxy_server.scheduler", scheduler): + await guard._check_file_deletion_allowed(file_id) + + @pytest.mark.asyncio + async def test_winning_pod_claims_the_row_between_fetching_and_billing(self): + """The claim flips batch_processed false -> true after the results are in hand and + before the spend log is written, so a concurrent pod's claim finds no matching row.""" + row = _FakeManagedObjectRow() + journal = [] + prisma = self._prisma(row, journal) + + with self._billing_patches(journal) as logging_obj: + await self._instance(prisma, self._router()).check_batch_cost() + + assert journal == ["fetch", "claim", "bill", "finalize"] + assert self._claim_calls(prisma) == [ + { + "where": {"id": "job-claim-1", "batch_processed": False}, + "data": {"batch_processed": True}, + } + ] + logging_obj.async_success_handler.assert_awaited_once() + assert row.batch_processed is True + + @pytest.mark.asyncio + async def test_a_pod_that_loses_the_claim_after_fetching_does_not_bill(self): + """Both pods select the row and fetch its results in the same window. The one whose + compare-and-swap finds the row already taken must not write a second spend log.""" + row = _FakeManagedObjectRow() + journal = [] + prisma = self._prisma(row, journal) + + async def _other_pod_wins_the_row(): + row.batch_processed = True + + with self._billing_patches(journal, during_fetch=_other_pod_wins_the_row) as logging_obj: + await self._instance(prisma, self._router()).check_batch_cost() + + assert journal == ["fetch"] + logging_obj.async_success_handler.assert_not_awaited() + assert self._claim_calls(prisma) == [ + { + "where": {"id": "job-claim-1", "batch_processed": False}, + "data": {"batch_processed": True}, + } + ] + + @pytest.mark.asyncio + async def test_a_failed_spend_log_write_releases_the_claim(self): + """A transient failure while billing a claimed batch must hand the row back, or its + spend is silently lost instead of being retried on the next cycle.""" + row = _FakeManagedObjectRow() + journal = [] + prisma = self._prisma(row, journal) + + with self._billing_patches(journal, bill_error=Exception("spend log write failed")): + await self._instance(prisma, self._router()).check_batch_cost() + + assert journal == ["fetch", "claim", "bill", "release"] + assert row.batch_processed is False + assert self._claim_calls(prisma)[-1] == { + "where": {"id": "job-claim-1", "batch_processed": True}, + "data": {"batch_processed": False}, + } + + @pytest.mark.asyncio + async def test_a_worker_interrupted_mid_costing_leaves_the_batch_billable(self): + """A pod killed while reading a batch's results must leave the row for a later + cycle. Claiming before the fetch marked the batch processed for good, so the pod + that died took that batch's spend with it and no other pod ever selected it.""" + row = _FakeManagedObjectRow() + journal = [] + prisma = self._prisma(row, journal) + reached_fetch = asyncio.Event() + + async def _never_returns(): + reached_fetch.set() + await asyncio.Event().wait() + + with self._billing_patches(journal, during_fetch=_never_returns) as logging_obj: + interrupted = asyncio.create_task( + self._instance(prisma, self._router()).check_batch_cost() + ) + await asyncio.wait_for(reached_fetch.wait(), timeout=5) + assert row.batch_processed is False, "an in-flight costing must not mark the row processed" + interrupted.cancel() + with pytest.raises(asyncio.CancelledError): + await interrupted + + assert journal == ["fetch"] + logging_obj.async_success_handler.assert_not_awaited() + + survivor_journal = [] + survivor_prisma = self._prisma(row, survivor_journal) + with self._billing_patches(survivor_journal) as survivor_logging: + await self._instance(survivor_prisma, self._router()).check_batch_cost() + + assert survivor_journal == ["fetch", "claim", "bill", "finalize"] + survivor_logging.async_success_handler.assert_awaited_once() + assert row.batch_processed is True + + @pytest.mark.asyncio + async def test_costing_in_flight_keeps_the_referenced_file_undeletable(self): + """The deletion guard only holds files whose batch still has batch_processed false, + so claiming the row before the fetch let a concurrent delete remove the very output + file the in-flight costing was about to read.""" + row = _FakeManagedObjectRow() + journal = [] + prisma = self._prisma(row, journal) + reached_fetch = asyncio.Event() + finish_fetch = asyncio.Event() + + async def _wait_for_the_delete_attempt(): + reached_fetch.set() + await finish_fetch.wait() + + with self._billing_patches(journal, during_fetch=_wait_for_the_delete_attempt): + costing = asyncio.create_task( + self._instance(prisma, self._router()).check_batch_cost() + ) + await asyncio.wait_for(reached_fetch.wait(), timeout=5) + + with pytest.raises(HTTPException) as blocked: + await self._run_deletion_guard(prisma, _CLAIM_OUTPUT_FILE_ID) + assert blocked.value.status_code == 400 + assert _CLAIM_OUTPUT_FILE_ID in blocked.value.detail + + finish_fetch.set() + await asyncio.wait_for(costing, timeout=5) + + assert journal == ["fetch", "claim", "bill", "finalize"] + assert row.batch_processed is True + await self._run_deletion_guard(prisma, _CLAIM_OUTPUT_FILE_ID) + + @pytest.mark.asyncio + async def test_schema_without_batch_processed_still_bills(self): + """Older schemas have no column to claim, so they keep the pre-fix behavior instead + of losing every batch's cost.""" + row = _FakeManagedObjectRow() + journal = [] + prisma = self._prisma(row, journal) + instance = self._instance(prisma, self._router()) + instance._has_batch_processed_column = False + + with self._billing_patches(journal) as logging_obj: + await instance.check_batch_cost() + + assert self._claim_calls(prisma) == [] + assert journal == ["fetch", "bill", "finalize"] + logging_obj.async_success_handler.assert_awaited_once() diff --git a/tests/proxy_unit_tests/test_custom_callback_input.py b/tests/proxy_unit_tests/test_custom_callback_input.py index 71a7e94b180..a032b8706bc 100644 --- a/tests/proxy_unit_tests/test_custom_callback_input.py +++ b/tests/proxy_unit_tests/test_custom_callback_input.py @@ -2,6 +2,7 @@ ## This test asserts the type of data passed into each method of the custom callback handler import asyncio import inspect +import json import os import sys import time diff --git a/tests/test_end_users.py b/tests/test_end_users.py index ff3cc4ec94b..bc1fcbb662d 100644 --- a/tests/test_end_users.py +++ b/tests/test_end_users.py @@ -14,47 +14,6 @@ from typing import Optional """ -async def chat_completion_with_headers(session, key, model="gpt-4"): - url = "http://0.0.0.0:4000/chat/completions" - headers = { - "Authorization": f"Bearer {key}", - "Content-Type": "application/json", - } - data = { - "model": model, - "messages": [ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Hello!"}, - ], - } - - async with session.post(url, headers=headers, json=data) as response: - status = response.status - response_text = await response.text() - - print(response_text) - print() - - if status != 200: - raise Exception(f"Request did not return a 200 status code: {status}") - - response_header_check( - response - ) # calling the function to check response headers - - raw_headers = response.raw_headers - raw_headers_json = {} - - for ( - item - ) in ( - response.raw_headers - ): # ((b'date', b'Fri, 19 Apr 2024 21:17:29 GMT'), (), ) - raw_headers_json[item[0].decode("utf-8")] = item[1].decode("utf-8") - - return raw_headers_json - - async def generate_key( session, i, diff --git a/tests/test_litellm/a2a_protocol/test_send_message_response.py b/tests/test_litellm/a2a_protocol/test_send_message_response.py index 832aa288c7a..ade7c72fc2e 100644 --- a/tests/test_litellm/a2a_protocol/test_send_message_response.py +++ b/tests/test_litellm/a2a_protocol/test_send_message_response.py @@ -32,12 +32,102 @@ def test_from_dict_preserves_existing_id(): assert response.id == "upstream-id" -def test_from_dict_without_request_id_still_requires_id(): - try: - LiteLLMSendMessageResponse.from_dict( - {"jsonrpc": "2.0", "error": {"code": -32054, "message": "x"}} - ) - except Exception as exc: - assert "id" in str(exc).lower() - else: - raise AssertionError("expected validation error when id and request_id missing") +def test_from_dict_preserves_integer_id_echoed_by_upstream(): + """JSON-RPC 2.0 types ``id`` as string|integer|null, and pydantic v2 does not + coerce int to str, so a str-only annotation rejects an upstream agent that + echoes an integer id. The value AND the type must survive.""" + payload = { + "id": 42, + "jsonrpc": "2.0", + "result": {"kind": "task"}, + } + + response = LiteLLMSendMessageResponse.from_dict(payload, request_id="r1") + + assert response.id == 42 + assert isinstance(response.id, int) + + +def test_from_dict_preserves_falsy_integer_id(): + """``0`` is a legal JSON-RPC id and is falsy, so it must not be mistaken for an + absent id and backfilled from the request id.""" + payload = {"id": 0, "jsonrpc": "2.0", "result": {}} + + response = LiteLLMSendMessageResponse.from_dict(payload, request_id="r1") + + assert response.id == 0 + + +def test_backfilled_id_keeps_the_request_id_type(): + """The proxy's A2A endpoint reads the caller's ``id`` straight off the request + body, so it can be an integer. JSON-RPC requires the response id to equal the + request id, so backfilling an omitted id must not stringify it: a caller that + sent ``7`` cannot correlate a response carrying ``"7"``. One test, both + directions, so neither can regress unnoticed.""" + agent_error = { + "jsonrpc": "2.0", + "error": {"code": -32054, "message": "Session not found"}, + } + + from_int = LiteLLMSendMessageResponse.from_dict(agent_error, request_id=7) + from_str = LiteLLMSendMessageResponse.from_dict(agent_error, request_id="7") + + assert from_int.id == 7 + assert isinstance(from_int.id, int) + assert from_str.id == "7" + assert isinstance(from_str.id, str) + + +def test_from_dict_accepts_null_id_when_the_error_cannot_be_correlated(): + """JSON-RPC 2.0 section 5 requires ``id`` to be null on an error that cannot be + matched to a request, which is exactly the case where the caller supplied no id + for the backfill to use. Rejecting it turned an agent's error into a proxy 500.""" + response = LiteLLMSendMessageResponse.from_dict( + {"jsonrpc": "2.0", "error": {"code": -32054, "message": "x"}} + ) + + assert response.id is None + assert response.error == {"code": -32054, "message": "x"} + + +def test_from_dict_accepts_null_id_echoed_by_upstream(): + """An agent may answer an uncorrelatable request with an explicit ``"id": null``. + That is a well-formed response, not a validation failure.""" + response = LiteLLMSendMessageResponse.from_dict( + {"id": None, "jsonrpc": "2.0", "error": {"code": -32600, "message": "bad"}} + ) + + assert response.id is None + + +def test_id_accepts_every_member_of_the_json_rpc_union_and_nothing_else(): + """One test pinning the whole ``string | integer | null`` union the spec defines, + so widening the annotation cannot silently become "accept anything".""" + for accepted in ("s1", 42, 0, None): + assert LiteLLMSendMessageResponse(id=accepted).id == accepted + + # ``True``/``False`` are in here because bool subclasses int: a non-strict integer + # half would accept them and relay them as 1/0. Direct construction bypasses + # normalization, so the model has to hold this line on its own. + for rejected in (True, False, 1.5, ["a"], {"a": 1}): + try: + LiteLLMSendMessageResponse(id=rejected) + except Exception: + continue + raise AssertionError(f"id={rejected!r} is outside the JSON-RPC union and must be rejected") + + +def test_boolean_id_is_never_relayed_as_an_integer(): + """``bool`` subclasses ``int``, so widening the annotation to accept integers also + made pydantic coerce a boolean id to 1 or 0. That is worse than rejecting it: an id + of ``1`` collides with a real integer id another in-flight request may be using. + Both directions in one test, since either alone leaves the other free to regress.""" + agent_error = {"jsonrpc": "2.0", "error": {"code": -32054, "message": "x"}} + + echoed = LiteLLMSendMessageResponse.from_dict({"id": True, "jsonrpc": "2.0", "result": {}}) + backfilled = LiteLLMSendMessageResponse.from_dict(agent_error, request_id=True) + + assert echoed.id == "True" + assert backfilled.id == "True" + assert not isinstance(echoed.id, int) + assert not isinstance(backfilled.id, int) diff --git a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py index 8ecb7f4c6f0..42b5ba235bc 100644 --- a/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py +++ b/tests/test_litellm/completion_extras/litellm_responses_transformation/test_completion_extras_litellm_responses_transformation_handler.py @@ -7,11 +7,13 @@ import pytest sys.path.insert(0, os.path.abspath("../../..")) +import litellm from litellm.completion_extras.litellm_responses_transformation.handler import ( ResponsesToCompletionBridgeHandler, ) from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging from litellm.litellm_core_utils.streaming_handler import CustomStreamWrapper +from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ModelResponse @@ -265,3 +267,74 @@ def test_completion_streams_completed_model_response(): assert "".join(chunk.choices[0].delta.content or "" for chunk in chunks) == "pong", ( f"completed response did not stream its content: {chunks}" ) + + +_PROVIDER_NATIVE_MODEL_CASES = [ + ("perplexity", "perplexity/kimi-k3", "perplexity/kimi-k3"), + ("perplexity", "openai/gpt-5.2", "openai/gpt-5.2"), + ("openai", "gpt-5.4", "gpt-5.4"), +] + + +def _upstream_model_for(handed_model: str, custom_llm_provider: str) -> str: + upstream_model, _, _, _ = litellm.get_llm_provider( + model=handed_model, + litellm_params=GenericLiteLLMParams(custom_llm_provider=custom_llm_provider), + ) + return upstream_model + + +@pytest.mark.parametrize( + "custom_llm_provider, bridge_model, expected_upstream_model", + _PROVIDER_NATIVE_MODEL_CASES, +) +def test_completion_keeps_provider_native_model_id_through_responses( + custom_llm_provider, bridge_model, expected_upstream_model +): + """responses() resolves the provider itself, so the bridge must not hand it an already-stripped model.""" + cached = ModelResponse(id="chatcmpl-cached", model=bridge_model) + bridge = ResponsesToCompletionBridgeHandler() + kwargs = _bridge_kwargs(stream=False) + kwargs["model"] = bridge_model + kwargs["custom_llm_provider"] = custom_llm_provider + + with ( + patch.object( + bridge.transformation_handler, + "transform_request", + return_value={"model": bridge_model, "input": "hi"}, + ), + patch("litellm.responses", return_value=cached) as responses_call, + ): + bridge.completion(**kwargs) + + handed_model = responses_call.call_args.kwargs["model"] + assert _upstream_model_for(handed_model, custom_llm_provider) == expected_upstream_model + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "custom_llm_provider, bridge_model, expected_upstream_model", + _PROVIDER_NATIVE_MODEL_CASES, +) +async def test_acompletion_keeps_provider_native_model_id_through_responses( + custom_llm_provider, bridge_model, expected_upstream_model +): + cached = ModelResponse(id="chatcmpl-cached-async", model=bridge_model) + bridge = ResponsesToCompletionBridgeHandler() + kwargs = _bridge_kwargs(stream=False) + kwargs["model"] = bridge_model + kwargs["custom_llm_provider"] = custom_llm_provider + + with ( + patch.object( + bridge.transformation_handler, + "transform_request", + return_value={"model": bridge_model, "input": "hi"}, + ), + patch("litellm.aresponses", new=AsyncMock(return_value=cached)) as responses_call, + ): + await bridge.acompletion(**kwargs) + + handed_model = responses_call.call_args.kwargs["model"] + assert _upstream_model_for(handed_model, custom_llm_provider) == expected_upstream_model diff --git a/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py b/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py index d9a0b275392..c75c8099ea1 100644 --- a/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py +++ b/tests/test_litellm/enterprise/proxy/test_managed_files_access_check.py @@ -192,6 +192,7 @@ async def test_check_batch_cost_should_call_afile_content_directly_with_credenti return_value=[mock_job] ) mock_prisma.db.litellm_managedobjecttable.update = AsyncMock() + mock_prisma.db.litellm_managedobjecttable.update_many = AsyncMock(return_value=1) # Mock proxy_logging_obj — should NOT be called for file content mock_proxy_logging = MagicMock() diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 7beb1c43a94..6c3c852395b 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -1,4 +1,5 @@ import asyncio +import base64 import os import sys from unittest.mock import AsyncMock, MagicMock, patch @@ -27,11 +28,16 @@ from litellm.experimental_mcp_client.client import ( MCPClient, _as_read_timeout, _first_non_cancelled_cause, + strip_auth_scheme, ) from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( classify_list_exception, list_fault_http_status, ) +from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( + _format_byok_openapi_auth_header, +) +from litellm.types.mcp_server.mcp_server_manager import MCPServer from litellm.types.mcp import MCPAuth, MCPStdioConfig, MCPTransport @@ -887,3 +893,159 @@ async def test_read_timeout_logs_an_actionable_line_that_quiet_on_error_cannot_d assert timeout_lines, f"expected an actionable timeout warning, got {warnings}" assert "http://upstream.local/mcp" in timeout_lines[0], "the line must name the server that stopped answering" assert "0.5s" in timeout_lines[0], "the line must name the budget that elapsed" + + +class TestAuthSchemeNormalization: + """MCP egress must emit exactly one authorization scheme. + + Callers supply both a bare credential and a complete header value (the latter whenever it is + passed through from ``x-mcp-auth`` / ``Authorization``), and the second shape used to be given + a second scheme, which upstream servers reject as a malformed token. + """ + + @pytest.mark.parametrize( + "auth_type, auth_value", + [ + (MCPAuth.bearer_token, "bare-token"), + (MCPAuth.bearer_token, "Bearer bare-token"), + (MCPAuth.bearer_token, "bearer bare-token"), + (MCPAuth.bearer_token, " BEARER bare-token"), + (MCPAuth.oauth2, "bare-token"), + (MCPAuth.oauth2, "Bearer bare-token"), + (MCPAuth.oauth2_token_exchange, "bare-token"), + (MCPAuth.oauth2_token_exchange, "Bearer bare-token"), + ], + ) + def test_bearer_family_emits_exactly_one_scheme(self, auth_type, auth_value): + client = MCPClient(server_url="http://example.com/mcp", auth_type=auth_type, auth_value=auth_value) + + assert client._get_auth_headers()["Authorization"] == "Bearer bare-token" + + @pytest.mark.parametrize("auth_value", ["bare-token", "token bare-token", "TOKEN bare-token"]) + def test_token_scheme_emits_exactly_one_scheme(self, auth_value): + client = MCPClient(server_url="http://example.com/mcp", auth_type=MCPAuth.token, auth_value=auth_value) + + assert client._get_auth_headers()["Authorization"] == "token bare-token" + + @pytest.mark.parametrize( + "auth_type, auth_value", + [ + (MCPAuth.bearer_token, "Bearertoken"), + (MCPAuth.oauth2, "Bearer.eyJzdWIiOiJhYmMifQ.sig"), + (MCPAuth.token, "tokenish"), + ], + ) + def test_a_credential_merely_starting_with_the_scheme_text_is_left_intact(self, auth_type, auth_value): + """RFC 7235 requires whitespace between scheme and credential, so a token whose first + characters happen to spell the scheme is a credential, not a schemed value.""" + client = MCPClient(server_url="http://example.com/mcp", auth_type=auth_type, auth_value=auth_value) + + scheme = "token" if auth_type == MCPAuth.token else "Bearer" + assert client._get_auth_headers()["Authorization"] == f"{scheme} {auth_value}" + + @pytest.mark.parametrize( + "auth_type, auth_value, expected", + [ + (MCPAuth.bearer_token, "Bearer ", "Bearer Bearer"), + (MCPAuth.bearer_token, "Bearer ", "Bearer Bearer"), + ], + ) + def test_a_scheme_with_no_credential_behind_it_still_produces_a_header(self, auth_type, auth_value, expected): + """Treating this as a schemed value would leave nothing to send, and a request with no + Authorization at all is harder to diagnose upstream than a visibly wrong one.""" + client = MCPClient(server_url="http://example.com/mcp", auth_type=auth_type, auth_value=auth_value) + + assert client._get_auth_headers()["Authorization"] == expected + + def test_basic_with_a_scheme_and_no_credential_still_produces_a_header(self): + client = MCPClient(server_url="http://example.com/mcp", auth_type=MCPAuth.basic, auth_value="Basic ") + + assert "Authorization" in client._get_auth_headers() + + def test_basic_accepts_an_already_encoded_schemed_value_without_re_encoding_it(self): + """Stripping the scheme at header-build time cannot fix this shape: ``to_basic_auth`` has by + then encoded the whole ``Basic ...`` string, leaving no prefix to find.""" + encoded = base64.b64encode(b"user:pass").decode() + + client = MCPClient( + server_url="http://example.com/mcp", + auth_type=MCPAuth.basic, + auth_value=f"Basic {encoded}", + ) + + header = client._get_auth_headers()["Authorization"] + assert header == f"Basic {encoded}" + assert base64.b64decode(header.split(" ", 1)[1]) == b"user:pass" + + @pytest.mark.parametrize("auth_value", ["user:pass", "Basic user:pass", "basic user:pass"]) + def test_basic_always_emits_encoded_credentials(self, auth_value): + """A schemed value whose remainder is raw rather than encoded is still a username/password + pair, so it is encoded rather than forwarded as an invalid RFC 7617 header.""" + client = MCPClient(server_url="http://example.com/mcp", auth_type=MCPAuth.basic, auth_value=auth_value) + + header = client._get_auth_headers()["Authorization"] + assert base64.b64decode(header.split(" ", 1)[1]) == b"user:pass" + + def test_authorization_auth_type_is_passed_through_verbatim(self): + """``MCPAuth.authorization`` means the caller owns the whole header value.""" + client = MCPClient( + server_url="http://example.com/mcp", + auth_type=MCPAuth.authorization, + auth_value="Bearer Bearer deliberately-doubled", + ) + + assert client._get_auth_headers()["Authorization"] == "Bearer Bearer deliberately-doubled" + + def test_api_key_credential_is_not_treated_as_a_schemed_value(self): + client = MCPClient( + server_url="http://example.com/mcp", + auth_type=MCPAuth.api_key, + auth_value="Bearer looks-schemed", + ) + + assert client._get_auth_headers()["X-API-Key"] == "Bearer looks-schemed" + + +@pytest.mark.parametrize( + "auth_value, scheme, expected", + [ + ("Bearer abc", "Bearer", "abc"), + ("bearer abc", "Bearer", "abc"), + (" Bearer abc ", "Bearer", "abc "), + ("abc", "Bearer", "abc"), + ("Bearerabc", "Bearer", "Bearerabc"), + ("Basic abc", "Bearer", "Basic abc"), + ("token abc", "token", "abc"), + ("Basic abc", "Basic", "abc"), + ("Bearer ", "Bearer", "Bearer "), + ("Bearer ", "Bearer", "Bearer "), + ], +) +def test_strip_auth_scheme(auth_value, scheme, expected): + assert strip_auth_scheme(auth_value, scheme) == expected + + +@pytest.mark.parametrize( + "auth_type, auth_value, expected", + [ + (MCPAuth.bearer_token, "Bearer jwt", "Bearer jwt"), + (MCPAuth.bearer_token, "jwt", "Bearer jwt"), + (MCPAuth.api_key, "ApiKey secret", "ApiKey secret"), + (MCPAuth.api_key, "secret", "ApiKey secret"), + (MCPAuth.basic, "Basic dXNlcjpwYXNz", "Basic dXNlcjpwYXNz"), + ], +) +def test_openapi_byok_auth_header_emits_exactly_one_scheme(auth_type, auth_value, expected): + """A non-BYOK server short-circuits ``_resolve_byok_mcp_auth_header``, so this formatter also + receives the deprecated global ``x-mcp-auth``, which is already a complete header value.""" + server = MCPServer( + server_id="s1", + name="openapi-server", + url="http://example.com/mcp", + transport=MCPTransport.http, + auth_type=auth_type, + spec_path="/tmp/spec.json", + ) + + assert server.is_byok is False + assert _format_byok_openapi_auth_header(server, auth_value) == expected diff --git a/tests/test_litellm/integrations/datadog/test_datadog_metrics.py b/tests/test_litellm/integrations/datadog/test_datadog_metrics.py index 2a26b7fade8..a4a4ca334b0 100644 --- a/tests/test_litellm/integrations/datadog/test_datadog_metrics.py +++ b/tests/test_litellm/integrations/datadog/test_datadog_metrics.py @@ -63,6 +63,38 @@ async def test_extract_tags(clean_env): assert "team:test-team" in tags +@pytest.mark.asyncio +async def test_extract_tags_normalizes_team_alias(clean_env): + """Team aliases with uppercase or special characters match what Datadog stores.""" + logger = DatadogMetricsLogger(start_periodic_flush=False) + + payload = StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4o", + metadata={"user_api_key_team_alias": "P&T CTO-B2B"}, + ) + + tags = logger._extract_tags(log=payload, status_code="200") + + assert "team:p_t_cto-b2b" in tags + + +@pytest.mark.asyncio +async def test_extract_tags_keeps_non_string_team_id(clean_env): + """A numeric team id still produces a team tag instead of aborting the metric.""" + logger = DatadogMetricsLogger(start_periodic_flush=False) + + payload = StandardLoggingPayload( + custom_llm_provider="openai", + model="gpt-4o", + metadata={"user_api_key_team_id": 67890}, + ) + + tags = logger._extract_tags(log=payload, status_code="200") + + assert "team:67890" in tags + + @pytest.mark.asyncio async def test_extract_tags_no_team(clean_env): """Test tag extraction when no team info is present.""" diff --git a/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py b/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py index cc9eae7a371..624995085aa 100644 --- a/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py +++ b/tests/test_litellm/integrations/datadog/test_datadog_tags_regression.py @@ -1,3 +1,4 @@ +import datetime import os import sys from unittest.mock import patch @@ -6,7 +7,8 @@ import pytest sys.path.insert(0, os.path.abspath("../../../")) -from litellm.integrations.datadog.datadog_handler import get_datadog_tags +from litellm.integrations.datadog.datadog import DataDogLogger +from litellm.integrations.datadog.datadog_handler import get_datadog_tags, normalize_datadog_tag_value from litellm.integrations.datadog.datadog_cost_management import ( DatadogCostManagementLogger, ) @@ -27,6 +29,7 @@ class TestDatadogTagsRegression: "POD_NAME": "test-pod", "DD_API_KEY": "mock-api-key", "DD_APP_KEY": "mock-app-key", + "DD_SITE": "test.datadoghq.com", }, ): yield @@ -58,6 +61,57 @@ class TestDatadogTagsRegression: # Verify NEW team tag is added assert "team:regression-team" in tags_with_team + @pytest.mark.parametrize( + ("value", "expected"), + ( + ("P&T", "p_t"), + ("CTO-B2B", "cto-b2b"), + (" Team & Key!! ", "team_key"), + ("regression-team", "regression-team"), + ), + ) + def test_normalize_datadog_tag_value(self, value, expected): + assert normalize_datadog_tag_value(value) == expected + + def test_get_datadog_tags_normalizes_alias_and_request_tag_values(self, mock_env_vars): + payload = StandardLoggingPayload( + request_tags=["capability:P&T"], + metadata=StandardLoggingMetadata(user_api_key_team_alias="CTO-B2B"), + ) + + tags = get_datadog_tags(payload) + + assert "request_tag:capability:p_t" in tags + assert "team:cto-b2b" in tags + + def test_get_datadog_tags_keeps_non_string_tag_values(self, mock_env_vars): + payload = StandardLoggingPayload( + request_tags=[12345, "capability:P&T"], + metadata=StandardLoggingMetadata(user_api_key_team_id=67890), + ) + + tags = get_datadog_tags(payload) + + assert "request_tag:12345" in tags + assert "request_tag:capability:p_t" in tags + assert "team:67890" in tags + + @pytest.mark.asyncio + async def test_non_string_request_tag_still_emits_the_datadog_payload(self, mock_env_vars): + with patch("asyncio.create_task"): + logger = DataDogLogger() + payload = StandardLoggingPayload(request_tags=[12345], metadata=StandardLoggingMetadata()) + + await logger.async_log_success_event( + kwargs={"standard_logging_object": payload}, + response_obj=None, + start_time=datetime.datetime(2026, 1, 1), + end_time=datetime.datetime(2026, 1, 1), + ) + + assert len(logger.log_queue) == 1 + assert "request_tag:12345" in logger.log_queue[0]["ddtags"].split(",") + @pytest.mark.asyncio async def test_datadog_cost_management_tags_regression(self, mock_env_vars): """ @@ -89,3 +143,32 @@ class TestDatadogTagsRegression: assert tags_new["env"] == "test-env" assert tags_new["user"] == "new-user" assert tags_new["team"] == "new-team-alias" # New feature verified + + @pytest.mark.asyncio + async def test_datadog_cost_management_normalizes_alias_and_custom_tag_values(self, mock_env_vars): + logger = DatadogCostManagementLogger(cost_tag_keys=["capability"]) + payload = StandardLoggingPayload( + request_tags=["capability:Space & Punctuation!"], + metadata=StandardLoggingMetadata( + user_api_key_alias="P&T", + user_api_key_team_alias="CTO-B2B", + ), + ) + + tags = logger._extract_tags(payload) + + assert tags["user"] == "p_t" + assert tags["team"] == "cto-b2b" + assert tags["capability"] == "space_punctuation" + + @pytest.mark.asyncio + async def test_datadog_cost_management_keeps_non_string_alias_values(self, mock_env_vars): + logger = DatadogCostManagementLogger() + payload = StandardLoggingPayload( + metadata=StandardLoggingMetadata(user_api_key_alias=12345, user_api_key_team_id=67890), + ) + + tags = logger._extract_tags(payload) + + assert tags["user"] == "12345" + assert tags["team"] == "67890" diff --git a/tests/test_litellm/integrations/gitlab/test_gitlab_integration.py b/tests/test_litellm/integrations/gitlab/test_gitlab_integration.py index 8a0ae030fff..8118af56b0e 100644 --- a/tests/test_litellm/integrations/gitlab/test_gitlab_integration.py +++ b/tests/test_litellm/integrations/gitlab/test_gitlab_integration.py @@ -92,7 +92,7 @@ def test_gitlab_prompt_manager_error_handling_load(mock_client_class): with pytest.raises( Exception, match="Failed to load prompt 'gitlab::oops' from GitLab" ): - GitLabPromptManager(config, prompt_id="oops").prompt_manager + _ = GitLabPromptManager(config, prompt_id="oops").prompt_manager def test_gitlab_prompt_manager_config_validation_via_client_ctor(): @@ -105,7 +105,7 @@ def test_gitlab_prompt_manager_config_validation_via_client_ctor(): side_effect=ValueError("project and access_token are required"), ): with pytest.raises(ValueError, match="project and access_token are required"): - GitLabPromptManager({}).prompt_manager + _ = GitLabPromptManager({}).prompt_manager # ----------------------------- diff --git a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py index e92a368d24b..7bf15f59eb9 100644 --- a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py +++ b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py @@ -1738,6 +1738,23 @@ class TestEnableAnthropicPromptCaching: assert result_msgs[-1]["content"][-1]["cache_control"] == {"type": "ephemeral"} assert "cache_control" not in result_msgs[0]["content"][-1] + def test_messages_with_default_injections_leaves_the_caller_list_untouched(self, monkeypatch): + """ + Routing calls this on the live request's own message list to derive the affinity key, before + the request is sent. Marking in place would leak litellm's breakpoints into the caller's + messages, where the real injection pass later reads them back as client-supplied ones. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + messages = copy.deepcopy(self.MESSAGES) + before = copy.deepcopy(messages) + + injected = AnthropicCacheControlHook.messages_with_default_injections( + messages=messages, models=("claude-sonnet-4-5",) + ) + + assert injected != messages + assert messages == before + class TestPerKeyEnablePromptCaching: """Per-request enable_prompt_caching override (stamped from key metadata) with the global flag off.""" diff --git a/tests/test_litellm/integrations/test_prometheus_metrics_endpoint.py b/tests/test_litellm/integrations/test_prometheus_metrics_endpoint.py new file mode 100644 index 00000000000..f0e6495ba22 --- /dev/null +++ b/tests/test_litellm/integrations/test_prometheus_metrics_endpoint.py @@ -0,0 +1,276 @@ +"""The /metrics app must render off the event loop, coalesce concurrent scrapes and stream chunks.""" + +from __future__ import annotations + +import asyncio +import threading +import time +from collections.abc import Iterator, Mapping, Sequence +from typing import Final + +import httpx +import pytest +from prometheus_client import CollectorRegistry, Gauge +from prometheus_client.metrics_core import GaugeMetricFamily +from prometheus_client.registry import Collector + +from litellm.integrations.prometheus_metrics_endpoint import ( + RESPONSE_CHUNK_SIZE_BYTES, + make_metrics_asgi_app, +) + +_GATE_TIMEOUT_SECONDS: Final = 10.0 +_SECOND_SCRAPE_SETTLE_SECONDS: Final = 0.2 + + +class _SlowCollector(Collector): + """Blocking collector standing in for a large registry render.""" + + def __init__(self, block_seconds: float, sample_count: int = 1) -> None: + self.block_seconds = block_seconds + self.sample_count = sample_count + self.collect_calls = 0 + + def collect(self) -> Iterator[GaugeMetricFamily]: + self.collect_calls += 1 + time.sleep(self.block_seconds) + family: Final = GaugeMetricFamily("slow_metric", "slow", labels=("idx",)) + for idx in range(self.sample_count): + family.add_metric((str(idx),), 1.0) + yield family + + +def _client(registry: CollectorRegistry) -> httpx.AsyncClient: + return httpx.AsyncClient( + transport=httpx.ASGITransport(app=make_metrics_asgi_app(registry)), + base_url="http://metrics.test", + ) + + +def _registry_with(collector: Collector) -> CollectorRegistry: + registry: Final = CollectorRegistry() + registry.register(collector) + return registry + + +async def _scrape(client: httpx.AsyncClient, headers: Mapping[str, str] | None = None) -> httpx.Response: + return await client.get("/metrics", headers=headers) + + +@pytest.mark.asyncio +async def test_render_does_not_block_the_event_loop(): + ticks: Final[list[float]] = [] # mutable-ok: records loop wakeups while the scrape is in flight + + async def ticker() -> None: + while True: + await asyncio.sleep(0.01) + ticks.append(time.monotonic()) + + ticker_task: Final = asyncio.create_task(ticker()) + async with _client(_registry_with(_SlowCollector(block_seconds=0.5))) as client: + try: + response: Final = await _scrape(client) + finally: + ticker_task.cancel() + + assert b"slow_metric" in response.content + assert len(ticks) > 5, "event loop was blocked while the registry was rendered" + + +@pytest.mark.asyncio +async def test_concurrent_identical_scrapes_share_one_render(): + collector: Final = _SlowCollector(block_seconds=0.2) + async with _client(_registry_with(collector)) as client: + responses: Final[Sequence[httpx.Response]] = await asyncio.gather(*(_scrape(client) for _ in range(5))) + + assert collector.collect_calls == 1 + for response in responses: + assert b"slow_metric" in response.content + + +@pytest.mark.asyncio +async def test_sequential_scrapes_are_rendered_fresh(): + collector: Final = _SlowCollector(block_seconds=0.0) + async with _client(_registry_with(collector)) as client: + await _scrape(client) + await _scrape(client) + + assert collector.collect_calls == 2 + + +@pytest.mark.asyncio +async def test_gzip_is_used_when_the_scraper_accepts_it(): + registry: Final = CollectorRegistry() + Gauge("plain_metric", "plain", registry=registry).set(1) + + async with _client(registry) as client: + compressed: Final = await _scrape(client, headers={"accept-encoding": "gzip"}) + plain: Final = await _scrape(client, headers={"accept-encoding": "identity"}) + + assert compressed.headers["content-encoding"] == "gzip" + assert "content-encoding" not in plain.headers + assert compressed.content == plain.content + assert b"plain_metric" in plain.content + + +@pytest.mark.asyncio +async def test_name_filter_restricts_the_rendered_registry(): + registry: Final = CollectorRegistry() + Gauge("wanted_metric", "wanted", registry=registry).set(1) + Gauge("other_metric", "other", registry=registry).set(1) + + async with _client(registry) as client: + response: Final = await client.get("/metrics", params={"name[]": "wanted_metric"}) + + assert b"wanted_metric" in response.content + assert b"other_metric" not in response.content + + +@pytest.mark.asyncio +async def test_large_payload_is_streamed_in_chunks(): + registry: Final = _registry_with(_SlowCollector(block_seconds=0.0, sample_count=5000)) + chunk_sizes: Final[list[int]] = [] # mutable-ok: records the ASGI body parts the app emitted + + async def send(message: Mapping[str, object]) -> None: + if message["type"] == "http.response.body": + body = message["body"] + assert isinstance(body, bytes) + chunk_sizes.append(len(body)) + + incoming: Final = iter(({"type": "http.request", "body": b"", "more_body": False},)) + + async def receive() -> Mapping[str, object]: + request: Final = next(incoming, None) + if request is not None: + return request + await asyncio.Event().wait() + return {"type": "http.disconnect"} + + app: Final = make_metrics_asgi_app(registry) + await app( + { + "type": "http", + "method": "GET", + "path": "/metrics", + "headers": (), + "query_string": b"", + }, + receive, + send, + ) + + assert sum(chunk_sizes) > RESPONSE_CHUNK_SIZE_BYTES + assert len(chunk_sizes) > 2 + assert max(chunk_sizes) <= RESPONSE_CHUNK_SIZE_BYTES + + +class _GatedCollector(Collector): + """Blocking collector that parks in the worker thread until the test releases it.""" + + def __init__(self) -> None: + self.started = threading.Event() + self.release = threading.Event() + self._lock = threading.Lock() + self.collect_calls = 0 + + def collect(self) -> Iterator[GaugeMetricFamily]: + with self._lock: + self.collect_calls += 1 + self.started.set() + self.release.wait(timeout=_GATE_TIMEOUT_SECONDS) + family: Final = GaugeMetricFamily("gated_metric", "gated") + family.add_metric((), 1.0) + yield family + + +async def _scrape_pair_concurrently( + registry: CollectorRegistry, collector: _GatedCollector, headers: Sequence[Mapping[str, str]] +) -> Sequence[httpx.Response]: + """Issue the second scrape only once the first one's render is parked inside the worker thread.""" + async with _client(registry) as client: + try: + first: Final = asyncio.create_task(_scrape(client, headers=headers[0])) + assert await asyncio.to_thread(collector.started.wait, _GATE_TIMEOUT_SECONDS), "first render never started" + second: Final = asyncio.create_task(_scrape(client, headers=headers[1])) + await asyncio.sleep(_SECOND_SCRAPE_SETTLE_SECONDS) + collector.release.set() + return await asyncio.gather(first, second) + finally: + collector.release.set() + + +@pytest.mark.parametrize("reverse", (False, True), ids=("as-listed", "reversed")) +@pytest.mark.parametrize( + "spellings", + ( + ({"accept-encoding": "gzip"}, {"accept-encoding": "gzip, deflate"}), + ({"accept": "*/*"}, {"accept": "text/plain;version=0.0.4;q=0.5,*/*;q=0.1"}), + ), + ids=("accept-encoding", "accept"), +) +@pytest.mark.asyncio +async def test_header_spellings_with_the_same_output_share_one_render( + spellings: Sequence[Mapping[str, str]], reverse: bool +): + collector: Final = _GatedCollector() + ordered: Final = tuple(reversed(spellings)) if reverse else spellings + + responses: Final = await _scrape_pair_concurrently(_registry_with(collector), collector, ordered) + + assert collector.collect_calls == 1, "the second scrape rendered the registry again instead of joining the first" + for response in responses: + assert b"gated_metric" in response.content + + +@pytest.mark.asyncio +async def test_different_output_formats_are_rendered_separately(): + collector: Final = _GatedCollector() + + responses: Final = await _scrape_pair_concurrently( + _registry_with(collector), + collector, + ({"accept": "text/plain"}, {"accept": "application/openmetrics-text"}), + ) + + assert collector.collect_calls == 2, "scrapes wanting different exposition formats must not share a render" + assert responses[0].headers["content-type"] != responses[1].headers["content-type"] + + +@pytest.mark.asyncio +async def test_concurrent_gzip_and_plain_scrapes_each_get_their_own_encoding(): + collector: Final = _GatedCollector() + + responses: Final = await _scrape_pair_concurrently( + _registry_with(collector), + collector, + ({"accept-encoding": "gzip"}, {"accept-encoding": "identity"}), + ) + + assert collector.collect_calls == 2, "scrapes wanting different content encodings must not share a render" + assert responses[0].headers["content-encoding"] == "gzip" + assert "content-encoding" not in responses[1].headers + for response in responses: + assert b"gated_metric" in response.content + + +@pytest.mark.asyncio +async def test_a_finishing_render_does_not_evict_another_that_is_still_in_flight(): + collector: Final = _GatedCollector() + async with _client(_registry_with(collector)) as client: + try: + parked: Final = asyncio.create_task(_scrape(client)) + assert await asyncio.to_thread(collector.started.wait, _GATE_TIMEOUT_SECONDS), "first render never started" + + unrelated: Final = await client.get("/metrics", params={"name[]": "no_such_metric"}) + assert unrelated.status_code == 200 + + joiner: Final = asyncio.create_task(_scrape(client)) + await asyncio.sleep(_SECOND_SCRAPE_SETTLE_SECONDS) + collector.release.set() + responses: Final = await asyncio.gather(parked, joiner) + finally: + collector.release.set() + + assert collector.collect_calls == 1, "an unrelated render finishing evicted the render still in flight" + for response in responses: + assert b"gated_metric" in response.content diff --git a/tests/test_litellm/integrations/test_shadow_eval_logger.py b/tests/test_litellm/integrations/test_shadow_eval_logger.py index 514d5c6adca..4f6fea7b710 100644 --- a/tests/test_litellm/integrations/test_shadow_eval_logger.py +++ b/tests/test_litellm/integrations/test_shadow_eval_logger.py @@ -40,11 +40,19 @@ def _job(**overrides) -> ActiveShadowEvalJob: return ActiveShadowEvalJob(**{**defaults, **overrides}) -def _prisma(jobs=(), attempt_counts=()) -> MagicMock: +def _prisma(jobs=(), attempt_counts=(), attempt_costs=()) -> MagicMock: + costs = {job_id: {"judge_cost": judge, "shadow_cost": shadow} for job_id, judge, shadow in attempt_costs} prisma = MagicMock() prisma.db.litellm_shadowevaljob.find_many = AsyncMock(return_value=list(jobs)) prisma.db.litellm_shadowevalattempt.group_by = AsyncMock( - return_value=[{"job_id": job_id, "_count": {"_all": count}} for job_id, count in attempt_counts] + return_value=[ + { + "job_id": job_id, + "_count": {"_all": count}, + "_sum": costs.get(job_id, {"judge_cost": 0.0, "shadow_cost": 0.0}), + } + for job_id, count in attempt_counts + ] ) prisma.db.litellm_shadowevalattempt.create = AsyncMock() return prisma @@ -61,6 +69,7 @@ def _job_record(job: ActiveShadowEvalJob, api_key_id="key-hash") -> MagicMock: shadow_percentage=job.shadow_percentage, judge_model=job.judge_model, max_turns=job.max_turns, + max_budget=job.max_budget, ends_at=job.ends_at, ).items(): setattr(record, field, value) @@ -92,13 +101,32 @@ def _router(shadow_text="shadow answer", judge_json='{"preference": "A", "confid return router -def _logger(router=None, prisma=None, jobs=()) -> ShadowEvalLogger: +def _spend_counter(store=None): + """In-memory stand-in for the proxy's cross-pod spend counter: reads take the max of + the counter and the caller's fallback, exactly like get_current_spend does for a key + shape the reseed helpers do not know.""" + counter = store if store is not None else {} + + async def read(key, fallback_spend, max_budget): + return max(counter.get(key, 0.0), fallback_spend) + + async def write(key, cost): + counter[key] = counter.get(key, 0.0) + cost + + return counter, read, write + + +def _logger(router=None, prisma=None, jobs=(), counter_store=None) -> ShadowEvalLogger: cache = InMemoryCache(max_size_in_memory=4, default_ttl=60) + counter, read, write = _spend_counter(counter_store) logger = ShadowEvalLogger( router_provider=lambda: router, prisma_provider=lambda: prisma, jobs_cache=cache, + job_spend_reader=read, + job_spend_writer=write, ) + logger._test_counter = counter if jobs: cache.set_cache("shadow_eval:active_jobs", {"key-hash": tuple(jobs)}) return logger @@ -447,7 +475,9 @@ def test_failure_detail_names_the_raising_frame(): except TypeError as e: detail = _failure_detail(e) lineno = e.__traceback__.tb_lineno - assert detail == f"TypeError at test_shadow_eval_logger.py:{lineno}: 'tuple' object does not support item assignment" + assert ( + detail == f"TypeError at test_shadow_eval_logger.py:{lineno}: 'tuple' object does not support item assignment" + ) try: raise ValueError("p" * 5 * _MAX_ERROR_CHARS) @@ -456,6 +486,73 @@ def test_failure_detail_names_the_raising_frame(): assert "ValueError at test_shadow_eval_logger.py:" in truncated_row_error +def test_call_cost_prefers_the_billed_figure_over_the_public_price_map(monkeypatch): + """The router client stamps _hidden_params.response_cost from the deployment's own + pricing; the public map reads 0 for deployment-priced models, so budgets gated on it + would never close. The map is only the fallback for responses with no stamp.""" + import litellm as litellm_module + from litellm.integrations.shadow_eval_logger import _call_cost + + monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.005) + stamped = MagicMock() + stamped._hidden_params = {"response_cost": 0.04} + assert _call_cost(stamped) == 0.04 + + from litellm.types.utils import HiddenParams + + object_stamped = MagicMock() + object_stamped._hidden_params = HiddenParams(response_cost=0.03) + assert _call_cost(object_stamped) == 0.03 + + unstamped = MagicMock() + unstamped._hidden_params = {"response_cost": None} + assert _call_cost(unstamped) == 0.005 + assert _call_cost({"choices": []}) == 0.005 + + +@pytest.mark.asyncio +async def test_a_cold_or_reset_counter_degrades_to_the_fill_floor_not_zero(monkeypatch: pytest.MonkeyPatch): + """The design leans on one owner contract: for a spend:shadow_eval:* key (no DB + reseed by design), get_current_spend returns the caller's fill-sum fallback whenever + the counter reads lower. A reset counter therefore degrades to the <=10s-stale DB + sum, never to zero, so a Redis expiry cannot re-open a spent budget by a full cap.""" + from litellm.proxy import proxy_server + + counter_key = "spend:shadow_eval:job-cold-test" + monkeypatch.setattr(proxy_server, "prisma_client", None) + proxy_server.spend_counter_cache.in_memory_cache.set_cache(key=counter_key, value=0.05) + try: + assert ( + await proxy_server.get_current_spend(counter_key=counter_key, fallback_spend=0.42, max_budget=1.0) == 0.42 + ) + proxy_server.spend_counter_cache.in_memory_cache.delete_cache(key=counter_key) + assert ( + await proxy_server.get_current_spend(counter_key=counter_key, fallback_spend=0.42, max_budget=1.0) == 0.42 + ) + finally: + proxy_server.spend_counter_cache.in_memory_cache.delete_cache(key=counter_key) + + +@pytest.mark.asyncio +async def test_an_unverifiable_budget_skips_the_sample_instead_of_spending(): + """A raising spend read (fail-closed enforcement, or an owner bug) must skip the + sample before any provider call, never admit it on a guess.""" + + async def unverifiable(key, fallback_spend, max_budget): + raise RuntimeError("budget unverifiable") + + prisma = _prisma() + router = _router() + logger = _logger(router=router, prisma=prisma, jobs=(_job(max_budget=1.0),)) + logger._read_job_spend = unverifiable + + await logger.async_log_success_event(_success_kwargs(), RESPONSE, None, None) + await _drain(logger) + + router.acompletion.assert_not_called() + prisma.db.litellm_shadowevalattempt.create.assert_not_called() + + def test_judge_prompt_is_bounded_however_large_the_inputs(): prompt = _judge_user_prompt("c" * 200_000, "a" * 200_000, "b" * 200_000) assert len(prompt) < _MAX_JUDGE_PROMPT_CHARS + 100 @@ -491,6 +588,7 @@ class TestSuccessHookSkipChain: assert row["shadow_model"] == "cheap-model" assert row["confidence"] == 0.9 assert row["judge_cost"] == 0.005 + assert row["shadow_cost"] == 0.005 assert row["error"] is None assert prisma.db.litellm_shadowevaljob.find_many.await_count == 0 @@ -580,6 +678,7 @@ class TestSuccessHookSkipChain: ({}, {"ends_at": datetime.now(timezone.utc) - timedelta(seconds=1)}), ({}, {"attempts": 200}), ({}, {"attempts": 199, "max_turns": 200, "_starts": 1}), + ({}, {"max_budget": 0.10, "spend": 0.10}), ], ids=[ "internal-origin", @@ -590,6 +689,7 @@ class TestSuccessHookSkipChain: "past-end", "turn-budget-reached", "budget-consumed-by-started-tasks", + "spend-budget-reached", ], ) async def test_skip_paths_store_nothing(self, kwargs_mutation, job_mutation): @@ -617,6 +717,61 @@ class TestSuccessHookSkipChain: assert prisma.db.litellm_shadowevalattempt.create.await_count == 1 + async def test_completed_pipelines_hold_spend_budget_within_a_cache_generation(self, monkeypatch): + """An attempt's recorded cost lands in the spend counter immediately, so the + second sample is skipped before any provider call even though the cached fill + still reads spend 0.""" + import litellm as litellm_module + + monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.005) + prisma = _prisma() + logger = _logger(router=_router(), prisma=prisma, jobs=(_job(max_budget=0.009, spend=0.0),)) + + await logger.async_log_success_event(_success_kwargs(request_id="req-1"), RESPONSE, None, None) + await _drain(logger) + await logger.async_log_success_event(_success_kwargs(request_id="req-2"), RESPONSE, None, None) + await _drain(logger) + + assert prisma.db.litellm_shadowevalattempt.create.await_count == 1 + assert logger._test_counter["spend:shadow_eval:job-1"] == 0.01 + + async def test_a_sibling_pod_sees_spend_through_the_shared_counter(self, monkeypatch): + """Two pods share the cross-pod counter: once pod A's attempts spend the budget, + pod B skips before its shadow call even though pod B's cached fill reads 0.""" + import litellm as litellm_module + + monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.005) + shared = {} + prisma_a = _prisma() + pod_a = _logger( + router=_router(), prisma=prisma_a, jobs=(_job(max_budget=0.009, spend=0.0),), counter_store=shared + ) + router_b = _router() + prisma_b = _prisma() + pod_b = _logger( + router=router_b, prisma=prisma_b, jobs=(_job(max_budget=0.009, spend=0.0),), counter_store=shared + ) + + await pod_a.async_log_success_event(_success_kwargs(request_id="req-1"), RESPONSE, None, None) + await _drain(pod_a) + await pod_b.async_log_success_event(_success_kwargs(request_id="req-2"), RESPONSE, None, None) + await _drain(pod_b) + + assert prisma_a.db.litellm_shadowevalattempt.create.await_count == 1 + prisma_b.db.litellm_shadowevalattempt.create.assert_not_called() + router_b.acompletion.assert_not_called() + + async def test_legacy_jobs_without_a_spend_budget_sample_on_turns_alone(self): + """A pre-migration job carries max_budget None: recorded spend must never gate it, + only its own max_turns can.""" + prisma = _prisma() + logger = _logger(router=_router(), prisma=prisma, jobs=(_job(max_budget=None, spend=999.0, attempts=5),)) + + await logger.async_log_success_event(_success_kwargs(request_id="req-1"), RESPONSE, None, None) + await _drain(logger) + + assert prisma.db.litellm_shadowevalattempt.create.await_count == 1 + async def test_v1_messages_surface_forwards_identity_from_litellm_metadata(self): """/v1/messages stores identity in litellm_params.litellm_metadata, so the hook resolves the bucket through the shared helper; every surface forwards the same @@ -714,7 +869,7 @@ class TestActiveJobsCache: async def test_cache_refill_resets_the_starts_counter(self): job = _job() - prisma = _prisma(jobs=[_job_record(job)], attempt_counts=[("job-1", 7)]) + prisma = _prisma(jobs=[_job_record(job)], attempt_counts=[("job-1", 7)], attempt_costs=[("job-1", 0.02, 0.03)]) logger = ShadowEvalLogger( router_provider=lambda: None, prisma_provider=lambda: prisma, @@ -722,9 +877,11 @@ class TestActiveJobsCache: ) logger._job_starts = {"job-1": 5} - await logger._active_jobs() + jobs = await logger._active_jobs() assert logger._job_starts == {} + assert jobs["key-hash"][0].attempts == 7 + assert jobs["key-hash"][0].spend == 0.05 @pytest.mark.asyncio @@ -749,9 +906,9 @@ class TestShadowPipeline: async def test_over_budget_key_skips_before_any_call(self, monkeypatch: pytest.MonkeyPatch): """The gate delegates to the auth path's own budget owner, so an over-budget verdict there (BudgetExceededError) skips the shadow before any provider call.""" - import litellm.proxy.auth.auth_checks as auth_checks from litellm.exceptions import BudgetExceededError from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth import auth_checks monkeypatch.setattr( auth_checks, @@ -777,13 +934,18 @@ class TestShadowPipeline: prisma.db.litellm_shadowevalattempt.create.assert_not_called() @pytest.mark.parametrize( - "router_factory,expected_error,expected_cost", + "router_factory,expected_error,expected_cost,expected_shadow_cost", [ - (lambda: _failing_router(), "provider exploded", 0.0), - (lambda: _router(judge_json="I prefer response A, definitely"), "unparseable judge verdict", 0.007), - (lambda: _router(judge_json='{"preference": "'), "unparseable judge verdict", 0.007), - (lambda: _router(judge_json="{}"), "unparseable judge verdict", 0.007), - (lambda: _router(judge_json='{"preference": "A", "confidence": "0.8'), "unparseable judge verdict", 0.007), + (lambda: _failing_router(), "provider exploded", 0.0, 0.0), + (lambda: _router(judge_json="I prefer response A, definitely"), "unparseable judge verdict", 0.007, 0.007), + (lambda: _router(judge_json='{"preference": "'), "unparseable judge verdict", 0.007, 0.007), + (lambda: _router(judge_json="{}"), "unparseable judge verdict", 0.007, 0.007), + ( + lambda: _router(judge_json='{"preference": "A", "confidence": "0.8'), + "unparseable judge verdict", + 0.007, + 0.007, + ), ], ids=[ "shadow-call-fails", @@ -794,7 +956,7 @@ class TestShadowPipeline: ], ) async def test_failures_become_error_rows_and_keep_billed_judge_cost( - self, router_factory, expected_error, expected_cost, monkeypatch: pytest.MonkeyPatch + self, router_factory, expected_error, expected_cost, expected_shadow_cost, monkeypatch: pytest.MonkeyPatch ): import litellm as litellm_module @@ -818,6 +980,66 @@ class TestShadowPipeline: assert expected_error in row["error"] assert row["confidence"] is None assert row["judge_cost"] == expected_cost + assert row["shadow_cost"] == expected_shadow_cost + + async def test_an_empty_shadow_reply_still_bills_its_cost(self, monkeypatch: pytest.MonkeyPatch): + """A shadow call that returns no extractable text has still billed; pricing it at + zero would keep the dollar gate open while shadow calls keep charging the key.""" + import litellm as litellm_module + + monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.007) + prisma = _prisma() + logger = _logger(router=_router(shadow_text=""), prisma=prisma) + + await logger._run_shadow_eval( + job=_job(), + request_id="req-1", + messages=({"role": "user", "content": "hi"},), + real_text="real answer", + real_model="claude-opus", + control_tier=None, + shadow_params={}, + parent_metadata={}, + ) + + row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"] + assert row["outcome"] == "error" + assert "empty response" in row["error"] + assert row["shadow_cost"] == 0.007 + assert logger._test_counter["spend:shadow_eval:job-1"] == 0.007 + + async def test_a_pipeline_error_after_the_shadow_call_keeps_its_billed_cost(self, monkeypatch: pytest.MonkeyPatch): + """An unexpected error between the billed shadow call and the attempt write must + still record the shadow cost, or the per-key dollar gate undercounts forever.""" + import litellm as litellm_module + import litellm.integrations.shadow_eval_logger as shadow_eval_module + + monkeypatch.setattr(litellm_module, "completion_cost", lambda completion_response: 0.007) + + def explode(conversation, response_a, response_b): + raise RuntimeError("judge prompt build failed") + + monkeypatch.setattr(shadow_eval_module, "_judge_user_prompt", explode) + prisma = _prisma() + logger = _logger(router=_router(), prisma=prisma) + + await logger._run_shadow_eval( + job=_job(), + request_id="req-1", + messages=({"role": "user", "content": "hi"},), + real_text="real answer", + real_model="claude-opus", + control_tier=None, + shadow_params={}, + parent_metadata={}, + ) + + row = prisma.db.litellm_shadowevalattempt.create.call_args.kwargs["data"] + assert row["outcome"] == "error" + assert "pipeline error" in row["error"] + assert row["shadow_cost"] == 0.007 + assert row["judge_cost"] == 0.0 + assert logger._test_counter["spend:shadow_eval:job-1"] == 0.007 async def test_sub_calls_carry_identity_and_origin_but_never_parent_request_state(self): prisma = _prisma() @@ -918,9 +1140,7 @@ class TestDirection: router = _router() logger = _logger(router=router, prisma=prisma, jobs=(_reverse_job(),)) - await logger.async_log_success_event( - _success_kwargs(request_metadata=_routed_by()), RESPONSE, None, None - ) + await logger.async_log_success_event(_success_kwargs(request_metadata=_routed_by()), RESPONSE, None, None) await _drain(logger) assert router.acompletion.call_args_list[0].kwargs["model"] == "baseline-model" @@ -967,9 +1187,7 @@ class TestDirection: jobs=(_job(id="forward-job", router_name="other-router"), _reverse_job(id="reverse-job")), ) - await logger.async_log_success_event( - _success_kwargs(request_metadata=_routed_by()), RESPONSE, None, None - ) + await logger.async_log_success_event(_success_kwargs(request_metadata=_routed_by()), RESPONSE, None, None) await _drain(logger) rows = [call.kwargs["data"] for call in prisma.db.litellm_shadowevalattempt.create.call_args_list] diff --git a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py b/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py index fb4cb494bee..956da571d43 100644 --- a/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py +++ b/tests/test_litellm/litellm_core_utils/test_get_litellm_params.py @@ -215,3 +215,32 @@ class TestMetadataFallsBackToLitellmMetadata: assert result["metadata"] is not litellm_metadata result["metadata"].pop("trace_id") assert litellm_metadata == {"trace_id": "trace-1"} + + +class TestRustOptIn: + """`rust: true` is a litellm param, so it has to reach `litellm_params`. + + `all_litellm_params` keeps it out of the provider body; without it also + being carried into `litellm_params` the chat completions handlers cannot + see the opt-in and the Rust path is silently never taken. + """ + + def test_rust_is_an_optional_kwargs_key(self): + assert "rust" in _OPTIONAL_KWARGS_KEYS + + def test_rust_is_forwarded_from_completion_kwargs(self): + from litellm.litellm_core_utils.get_litellm_params import FORWARDED_KWARGS_KEYS + + assert "rust" in FORWARDED_KWARGS_KEYS + + def test_rust_survives_into_litellm_params(self): + params = get_litellm_params(rust=True) + assert params["rust"] is True + + def test_rust_is_absent_when_the_deployment_did_not_set_it(self): + assert "rust" not in get_litellm_params() + + def test_rust_stays_out_of_the_provider_body(self): + from litellm.types.utils import all_litellm_params + + assert "rust" in all_litellm_params diff --git a/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py b/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py index 270c59f595f..b953bfaa565 100644 --- a/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py +++ b/tests/test_litellm/litellm_core_utils/test_ptu_pricing.py @@ -7,6 +7,7 @@ from unittest.mock import patch import pytest from litellm.litellm_core_utils.ptu_pricing import ( + ptu_config_error, CUSTOM_PRICING_FIELDS, PTU_EMPTIED_PRICING_FIELDS, PTU_ZEROED_PRICING_FIELDS, @@ -161,3 +162,50 @@ def test_a_setting_that_is_not_a_charge_is_left_alone(): assert override is not None assert "output_vector_size" not in override + + +# --- the rule both the endpoints and config.yaml registration enforce --------------- + + +def test_a_complete_reservation_has_no_error(): + assert ptu_config_error(_VALID) is None + + +def test_a_deployment_with_no_ptu_fields_is_not_a_ptu_deployment(): + """The gate must stay scoped to PTU configuration, or it would reject every ordinary + deployment for lacking a team_id.""" + assert ptu_config_error({"team_id": "team-alpha"}) is None + assert ptu_config_error({}) is None + + +@pytest.mark.parametrize( + "override, expected", + [ + ({"team_id": None}, "team_id is required when PTU fields are set (one model maps to one team)"), + ({"team_id": ""}, "team_id is required when PTU fields are set (one model maps to one team)"), + ({"cost_per_ptu_per_hour": None}, "ptu_count and cost_per_ptu_per_hour must be set together"), + ({"ptu_count": None}, "ptu_count and cost_per_ptu_per_hour must be set together"), + ({"ptu_effective_to": "2025-01-01T00:00:00Z"}, "ptu_effective_to must be after ptu_effective_from"), + ], + ids=["no team", "blank team", "count without rate", "rate without count", "inverted window"], +) +def test_an_incoherent_reservation_names_its_reason(override, expected): + assert ptu_config_error({**_VALID, **override}) == expected + + +def test_a_missing_start_is_explained_rather_than_inferred(): + error = ptu_config_error({k: v for k, v in _VALID.items() if k != "ptu_effective_from"}) + + assert error is not None + assert error.startswith("ptu_effective_from is required when PTU fields are set") + + +def test_an_inverted_window_is_caught_before_the_count_and_rate_gate(): + """A patch that moves one end of the window carries no count or rate, so ordering has to + be checked first or an inverted window reaches the row and the next load cannot parse it.""" + window_only = { + "ptu_effective_from": "2026-01-01T00:00:00Z", + "ptu_effective_to": "2025-01-01T00:00:00Z", + } + + assert ptu_config_error(window_only) == "ptu_effective_to must be after ptu_effective_from" diff --git a/tests/test_litellm/litellm_core_utils/test_thread_pool_executor.py b/tests/test_litellm/litellm_core_utils/test_thread_pool_executor.py new file mode 100644 index 00000000000..e81de277eaa --- /dev/null +++ b/tests/test_litellm/litellm_core_utils/test_thread_pool_executor.py @@ -0,0 +1,128 @@ +import logging +import threading +import time +from typing import Final + +from litellm._logging import verbose_logger +from litellm.constants import LOGGING_EXECUTOR_MAX_PENDING_TASKS +from litellm.litellm_core_utils.thread_pool_executor import ( + BoundedLoggingThreadPoolExecutor, + executor, +) + + +def test_submit_drops_tasks_when_backlog_is_full(): + release: Final = threading.Event() + started: Final = threading.Event() + ran_first: Final = threading.Event() + ran_second: Final = threading.Event() + ran_dropped: Final = threading.Event() + + def blocking_task(ran: threading.Event) -> None: + ran.set() + started.set() + release.wait(timeout=10) + + pool: Final = BoundedLoggingThreadPoolExecutor(max_workers=1, max_pending_tasks=2) + try: + first: Final = pool.submit(blocking_task, ran_first) + assert started.wait(timeout=10) + second: Final = pool.submit(blocking_task, ran_second) + dropped: Final = pool.submit(blocking_task, ran_dropped) + + assert dropped.cancelled() + assert not first.cancelled() + assert not second.cancelled() + + release.set() + first.result(timeout=10) + second.result(timeout=10) + assert ran_first.is_set() + assert ran_second.is_set() + assert not ran_dropped.is_set() + finally: + release.set() + pool.shutdown(wait=True) + + +def test_submit_releases_slots_after_completion(): + pool: Final = BoundedLoggingThreadPoolExecutor(max_workers=1, max_pending_tasks=1) + + def submit_and_wait() -> str: + future: Final = pool.submit(lambda: "ok") + assert not future.cancelled() + return future.result(timeout=10) + + try: + results: Final = tuple(submit_and_wait() for _ in range(5)) + assert results == ("ok",) * 5 + finally: + pool.shutdown(wait=True) + + +def test_drop_warning_is_rate_limited(caplog): + release: Final = threading.Event() + started: Final = threading.Event() + + def blocking_task() -> None: + started.set() + release.wait(timeout=10) + + drop_logger: Final = logging.getLogger("test_bounded_logging_executor") + pool: Final = BoundedLoggingThreadPoolExecutor( + max_workers=1, + max_pending_tasks=1, + drop_log_interval_seconds=60.0, + logger=drop_logger, + ) + try: + pool.submit(blocking_task) + assert started.wait(timeout=10) + + with caplog.at_level(logging.WARNING, logger=drop_logger.name): + assert pool.submit(time.sleep, 0).cancelled() + assert pool.submit(time.sleep, 0).cancelled() + assert pool.submit(time.sleep, 0).cancelled() + + warnings: Final = tuple(record for record in caplog.records if record.name == drop_logger.name) + assert len(warnings) == 1 + assert warnings[0].args == (1, 1) + finally: + release.set() + pool.shutdown(wait=True) + + +def test_each_drop_warning_counts_only_drops_since_the_last_one(caplog): + release: Final = threading.Event() + started: Final = threading.Event() + + def blocking_task() -> None: + started.set() + release.wait(timeout=10) + + drop_logger: Final = logging.getLogger("test_bounded_logging_executor_every_drop") + pool: Final = BoundedLoggingThreadPoolExecutor( + max_workers=1, + max_pending_tasks=1, + drop_log_interval_seconds=0.0, + logger=drop_logger, + ) + try: + pool.submit(blocking_task) + assert started.wait(timeout=10) + + with caplog.at_level(logging.WARNING, logger=drop_logger.name): + assert pool.submit(time.sleep, 0).cancelled() + assert pool.submit(time.sleep, 0).cancelled() + + warnings: Final = tuple(record for record in caplog.records if record.name == drop_logger.name) + assert tuple(record.args for record in warnings) == ((1, 1), (1, 1)) + finally: + release.set() + pool.shutdown(wait=True) + + +def test_global_executor_is_bounded(): + assert isinstance(executor, BoundedLoggingThreadPoolExecutor) + assert executor._max_pending_tasks == LOGGING_EXECUTOR_MAX_PENDING_TASKS + assert executor._logger is verbose_logger diff --git a/tests/test_litellm/litellm_core_utils/test_token_counter.py b/tests/test_litellm/litellm_core_utils/test_token_counter.py index 71e686563a5..3c33ee13c3f 100644 --- a/tests/test_litellm/litellm_core_utils/test_token_counter.py +++ b/tests/test_litellm/litellm_core_utils/test_token_counter.py @@ -1,5 +1,6 @@ #### What this tests #### # This tests litellm.token_counter.token_counter() function +import importlib import os import sys import time @@ -7,6 +8,7 @@ import traceback from unittest.mock import MagicMock import pytest +import tiktoken sys.path.insert( 0, os.path.abspath("../../..") @@ -16,6 +18,8 @@ from unittest.mock import AsyncMock, MagicMock, patch import litellm from litellm import create_pretrained_tokenizer, decode, encode, get_modified_max_tokens from litellm import token_counter as token_counter_old +import litellm.constants +from litellm.litellm_core_utils.token_counter import _get_tiktoken_count_function from litellm.litellm_core_utils.token_counter import token_counter as token_counter_new from tests.large_text import text from tests.test_litellm.litellm_core_utils.messages_with_counts import ( @@ -54,6 +58,73 @@ def test_token_counter_basic(): ) +def test_token_counter_large_repeated_text_is_fast(): + messages = [{"role": "user", "content": [{"type": "text", "text": "A" * 1024 * 1024}]}] + + start_time = time.perf_counter() + tokens = token_counter_new(model="us.anthropic.claude-sonnet-4-6", messages=messages) + elapsed = time.perf_counter() - start_time + + assert elapsed < 2, f"Token counting took too long: {elapsed:.2f}s" + assert tokens > 0 + + +@pytest.mark.parametrize( + "text", + [ + "Short text", + "This is a normal message with punctuation, numbers, and a few words.", + ], +) +def test_token_counter_short_text_matches_tiktoken(text): + encoding = tiktoken.get_encoding("cl100k_base") + expected = len(encoding.encode(text, disallowed_special=())) + + assert token_counter_new(model="us.anthropic.claude-sonnet-4-6", text=text) == expected + + +def test_token_counter_text_over_chunk_boundary_stays_close_to_tiktoken(): + text = ("The quick brown fox jumps over the lazy dog. " * 30)[:1025] + encoding = tiktoken.get_encoding("cl100k_base") + expected = len(encoding.encode(text, disallowed_special=())) + + actual = token_counter_new(model="us.anthropic.claude-sonnet-4-6", text=text) + + assert abs(actual - expected) <= 4 + + +@pytest.mark.parametrize( + "configured", + ["0", "-1", "-1024", "not-an-int", "", " ", "999999999", "inf", "1e9"], +) +def test_invalid_chunk_size_config_stays_usable(monkeypatch, configured): + """A misconfigured chunk size must not raise, count zero, or restore the quadratic encode cost.""" + monkeypatch.setenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS", configured) + try: + reloaded = importlib.reload(litellm.constants) + chunk_size = reloaded.TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS + assert 1 <= chunk_size <= reloaded.TIKTOKEN_ENCODE_MAX_CHUNK_SIZE_CHARS + + encoding = tiktoken.get_encoding("cl100k_base") + count_tokens = _get_tiktoken_count_function( + lambda text: len(encoding.encode(text, disallowed_special=())), + chunk_size=chunk_size, + ) + assert count_tokens("The quick brown fox jumps over the lazy dog. " * 40) > 0 + finally: + monkeypatch.delenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS") + importlib.reload(litellm.constants) + + +def test_valid_chunk_size_config_is_honoured(monkeypatch): + monkeypatch.setenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS", "2048") + try: + assert importlib.reload(litellm.constants).TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS == 2048 + finally: + monkeypatch.delenv("TIKTOKEN_ENCODE_CHUNK_SIZE_CHARS") + importlib.reload(litellm.constants) + + def test_token_counter_with_prefix(): messages = [ {"role": "user", "content": "Who won the world cup in 2022?"}, @@ -974,7 +1045,7 @@ def test_token_counter_with_image_url(): try: token_counter(model="gpt-3.5-turbo", messages=messages_invalid) - assert False, "Expected ValueError for invalid detail value" + pytest.fail("Expected ValueError for invalid detail value") except ValueError as e: assert "Invalid detail value" in str( e diff --git a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py index f934c7184f8..f6cd6ac6734 100644 --- a/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py +++ b/tests/test_litellm/llms/anthropic/chat/test_anthropic_chat_handler.py @@ -1,7 +1,7 @@ import json import threading from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer -from unittest.mock import AsyncMock, MagicMock +from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -2045,3 +2045,389 @@ def test_non_bash_tool_result_skipped(): assert ( len(code_results) == 0 ), f"Expected 0 code_interpreter_results for text_editor result, got {len(code_results)}" + + +class TestRustChatCompletionsHook: + """The `rust: true` opt-in on `/chat/completions` for the Anthropic provider. + + The native callables are dependency-injected, so these run without the + compiled extension. + """ + + RUST_RESPONSE = { + "created": 1_700_000_000, + "model": "claude-sonnet-4-5-20260101", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hello from rust"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 11, + "completion_tokens": 4, + "total_tokens": 15, + "prompt_tokens_details": { + "cached_tokens": 0, + "cache_creation_tokens": 0, + "text_tokens": 11, + }, + }, + } + + @pytest.fixture(autouse=True) + def _reset_bridge(self, monkeypatch): + from litellm.rust_bridge import chat_completions as bridge + + monkeypatch.delenv("LITELLM_RUST", raising=False) + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) + yield + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) + + @staticmethod + def _completion_kwargs(**overrides): + from litellm.types.utils import ModelResponse + + kwargs = { + "model": "claude-sonnet-4-5", + "messages": [{"role": "user", "content": "hi"}], + "api_base": "https://api.anthropic.com/v1/messages", + "custom_llm_provider": "anthropic", + "custom_prompt_dict": {}, + "model_response": ModelResponse(), + "print_verbose": lambda *_args, **_kwargs: None, + "encoding": None, + "api_key": "sk-ant-test", + "logging_obj": MagicMock(), + "optional_params": {"max_tokens": 16}, + "timeout": 30.0, + "litellm_params": {"rust": True}, + "acompletion": False, + "headers": {}, + "client": None, + } + kwargs.update(overrides) + return kwargs + + @staticmethod + def _recording_logging_obj(): + """A logging object that keeps each hook's payload in a real list, so a + test can assert which path logged and what it carried.""" + calls = {"pre_call": [], "post_call": []} + logging_obj = MagicMock() + logging_obj.pre_call.side_effect = lambda **kwargs: calls["pre_call"].append(kwargs) + logging_obj.post_call.side_effect = lambda **kwargs: calls["post_call"].append(kwargs) + return logging_obj, calls + + def _inject(self, *, decline_reason=None, sync_result=None, sync_error=None): + from litellm.rust_bridge import chat_completions as bridge + + seen = {"gate": [], "call": []} + + def gate(**kwargs): + seen["gate"].append(kwargs) + return decline_reason + + def native(**kwargs): + seen["call"].append(kwargs) + if sync_error is not None: + raise sync_error + return dict(sync_result if sync_result is not None else self.RUST_RESPONSE) + + bridge.set_rust_chat_completions(decline=gate, chat_completions=native) + return seen + + def test_rust_true_serves_the_call_and_stamps_the_header(self): + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + + seen = self._inject() + response = AnthropicChatCompletion().completion(**self._completion_kwargs()) + + assert response.choices[0].message.content == "hello from rust" + assert response._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + assert len(seen["call"]) == 1 + + def test_the_core_receives_the_untranslated_openai_messages(self): + """Rust owns the translation, so the handler must not pre-translate.""" + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + + seen = self._inject() + AnthropicChatCompletion().completion( + **self._completion_kwargs( + messages=[ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + ] + ) + ) + assert seen["call"][0]["messages"] == [ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + ] + + def test_the_anthropic_max_tokens_default_is_merged_in_before_the_gate(self): + """`transform_request` applies `AnthropicConfig.get_config`; the Rust + path skips it, so the handler has to merge it or Anthropic 400s on a + request that omits `max_tokens`.""" + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + + seen = self._inject() + AnthropicChatCompletion().completion(**self._completion_kwargs(optional_params={})) + assert "max_tokens" in seen["gate"][0]["optional_params"] + assert seen["call"][0]["optional_params"]["max_tokens"] > 0 + + def test_a_caller_supplied_max_tokens_outranks_the_default(self): + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + + seen = self._inject() + AnthropicChatCompletion().completion( + **self._completion_kwargs(optional_params={"max_tokens": 7}) + ) + assert seen["call"][0]["optional_params"]["max_tokens"] == 7 + + def test_without_the_opt_in_the_core_is_never_consulted(self): + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + seen = self._inject() + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ) as transform, patch.object( + AnthropicChatCompletion, "acompletion_function" + ): + try: + AnthropicChatCompletion().completion( + **self._completion_kwargs(litellm_params={}) + ) + except Exception: + # The Python path goes on to make an HTTP call; reaching it is + # the assertion, so the network failure below is expected. + pass + assert seen["gate"] == [] + assert seen["call"] == [] + assert transform.called + + def test_a_declined_request_never_reaches_the_native_call(self): + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + seen = self._inject(decline_reason="unrecognized request parameter") + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): + try: + AnthropicChatCompletion().completion(**self._completion_kwargs()) + except Exception: + pass + assert len(seen["gate"]) == 1 + assert seen["call"] == [] + + def test_streaming_stays_on_the_python_path(self): + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + seen = self._inject() + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): + try: + AnthropicChatCompletion().completion( + **self._completion_kwargs(optional_params={"max_tokens": 16, "stream": True}) + ) + except Exception: + pass + assert seen["gate"] == [] + + def test_pre_call_logging_fires_exactly_once_on_the_rust_path(self): + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + + seen = self._inject() + logging_obj = MagicMock() + AnthropicChatCompletion().completion( + **self._completion_kwargs(logging_obj=logging_obj) + ) + assert logging_obj.pre_call.call_count == 1 + assert len(seen["call"]) == 1 + + def test_post_call_logging_fires_on_the_rust_path(self): + """The Rust core owns the provider call, so the Python transform that + normally raises `post_call` never runs. Without the bridge hook every + post_call callback goes silent and `original_response` stays unset.""" + import json + + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + + self._inject() + logging_obj = MagicMock() + AnthropicChatCompletion().completion( + **self._completion_kwargs(logging_obj=logging_obj) + ) + + assert logging_obj.post_call.call_count == 1 + logged = logging_obj.post_call.call_args.kwargs["original_response"] + assert json.loads(logged)["choices"][0]["message"]["content"] == "hello from rust" + + def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(self, monkeypatch): + """A decline never reached the provider, so the Python path serves the + request and owns the only post_call. Firing the hook there too would + double every post_call callback for one request.""" + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + from litellm.rust_bridge import chat_completions as bridge + + class _Declined(Exception): + pass + + class _FakeNative: + RustBridgeDeclined = _Declined + RustUpstreamError = type("_Upstream", (Exception,), {}) + + def declining_native(**_kwargs): + raise _Declined("blank message text") + + monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, chat_completions=declining_native + ) + + logging_obj, calls = self._recording_logging_obj() + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): + try: + AnthropicChatCompletion().completion( + **self._completion_kwargs(logging_obj=logging_obj) + ) + except Exception: + # The Python path goes on to make an HTTP call; the log count is + # the assertion, so a failure past this point is expected. + pass + + assert calls["post_call"] == [] + + @pytest.mark.asyncio + async def test_the_async_path_falls_back_when_the_core_declines(self, monkeypatch): + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + from litellm.rust_bridge import chat_completions as bridge + + class _Declined(Exception): + pass + + class _FakeNative: + RustBridgeDeclined = _Declined + RustUpstreamError = type("_Upstream", (Exception,), {}) + + monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + + async def declining_native(**_kwargs): + raise _Declined("blank message text") + + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=declining_native + ) + + sentinel = object() + + async def python_path(**_kwargs): + return sentinel + + with patch.object( + AnthropicChatCompletion, "acompletion_function", side_effect=python_path + ) as python_call: + result = await AnthropicChatCompletion().completion( + **self._completion_kwargs(acompletion=True) + ) + + assert result is sentinel + assert python_call.called, "a failing rust call must re-enter the python path" + + @pytest.mark.asyncio + async def test_the_async_path_serves_the_rust_response_without_the_fallback(self): + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + from litellm.rust_bridge import chat_completions as bridge + + async def native(**_kwargs): + return dict(self.RUST_RESPONSE) + + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=native + ) + + with patch.object(AnthropicChatCompletion, "acompletion_function") as python_call: + result = await AnthropicChatCompletion().completion( + **self._completion_kwargs(acompletion=True) + ) + + assert result.choices[0].message.content == "hello from rust" + assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + assert not python_call.called + + + def test_pre_call_logging_fires_once_when_the_sync_rust_call_declines(self, monkeypatch): + """One request, one pre_call, on the synchronous path too. Without the + suppression the Python path logs a second time for the same attempt.""" + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + from litellm.rust_bridge import chat_completions as bridge + + class _Declined(Exception): + pass + + class _FakeNative: + RustBridgeDeclined = _Declined + RustUpstreamError = type("_Upstream", (Exception,), {}) + + monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + + def declining_native(**_kwargs): + raise _Declined("blank message text") + + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, chat_completions=declining_native + ) + + logging_obj, calls = self._recording_logging_obj() + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): + try: + AnthropicChatCompletion().completion( + **self._completion_kwargs(logging_obj=logging_obj) + ) + except Exception: + # The Python path goes on to make an HTTP call; the log count is + # the assertion, so a failure past this point is expected. + pass + + assert len(calls["pre_call"]) == 1 + assert calls["pre_call"][0]["additional_args"]["complete_input_dict"]["model"] == ( + "claude-sonnet-4-5" + ) + + def test_pre_call_logging_still_fires_when_rust_is_not_involved(self, monkeypatch): + """The suppression must not swallow the log on the ordinary path.""" + from litellm.llms.anthropic.chat.handler import AnthropicChatCompletion + from litellm.llms.anthropic.chat.transformation import AnthropicConfig + + self._inject() + logging_obj, calls = self._recording_logging_obj() + with patch.object( + AnthropicConfig, "transform_request", return_value={"model": "m", "messages": []} + ): + try: + AnthropicChatCompletion().completion( + **self._completion_kwargs(litellm_params={}, logging_obj=logging_obj) + ) + except Exception: + pass + + assert len(calls["pre_call"]) == 1 + assert calls["pre_call"][0]["additional_args"]["complete_input_dict"] == { + "model": "m", + "messages": [], + } diff --git a/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py new file mode 100644 index 00000000000..8e67a7e3438 --- /dev/null +++ b/tests/test_litellm/llms/bedrock/chat/test_bedrock_converse_handler.py @@ -0,0 +1,489 @@ +"""Tests for `BedrockConverseLLM.completion`'s Rust chat completions hook. + +The native callables are dependency-injected, so these run without the compiled +extension, and AWS credential resolution is stubbed so nothing reaches STS. +""" + +from __future__ import annotations + +from unittest.mock import MagicMock, patch + +import httpx +import pytest + +from botocore.credentials import Credentials +from litellm.llms.bedrock.chat.converse_handler import BedrockConverseLLM +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.rust_bridge import chat_completions as bridge +from litellm.types.utils import ModelResponse + +RUST_RESPONSE = { + "created": 1_700_000_000, + "model": "anthropic.claude-sonnet-4-5-v1:0", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hello from rust"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 11, + "completion_tokens": 4, + "total_tokens": 15, + "prompt_tokens_details": { + "cached_tokens": 0, + "cache_creation_tokens": 0, + "text_tokens": 11, + }, + }, +} + +RESOLVED_CREDENTIALS = Credentials( + access_key="AKIARESOLVED", + secret_key="resolved-secret", + token="resolved-token", +) + + +@pytest.fixture(autouse=True) +def reset_bridge(monkeypatch): + monkeypatch.delenv("LITELLM_RUST", raising=False) + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) + yield + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) + + +def _inject(*, decline_reason=None, error: Exception | None = None): + seen: dict[str, list[dict]] = {"gate": [], "call": []} + + def gate(**kwargs): + seen["gate"].append(kwargs) + return decline_reason + + def native(**kwargs): + seen["call"].append(kwargs) + if error is not None: + raise error + return dict(RUST_RESPONSE) + + bridge.set_rust_chat_completions(decline=gate, chat_completions=native) + return seen + + +def _completion_kwargs(**overrides): + kwargs = { + "model": "bedrock/us-east-1/anthropic.claude-sonnet-4-5-v1:0", + "messages": [{"role": "user", "content": "hi"}], + "api_base": None, + "custom_prompt_dict": {}, + "model_response": ModelResponse(), + "encoding": None, + "logging_obj": MagicMock(), + "optional_params": {"maxTokens": 16}, + "acompletion": False, + "timeout": 30.0, + "litellm_params": {"rust": True}, + "extra_headers": None, + "client": None, + "api_key": None, + } + kwargs.update(overrides) + return kwargs + + +def _run(**overrides): + with patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ): + return BedrockConverseLLM().completion(**_completion_kwargs(**overrides)) + + +def _recording_logging_obj(): + """A logging object that keeps each hook's payload in a real list, so a test + can assert which path logged and what it carried.""" + calls = {"pre_call": [], "post_call": []} + logging_obj = MagicMock() + logging_obj.pre_call.side_effect = lambda **kwargs: calls["pre_call"].append(kwargs) + logging_obj.post_call.side_effect = lambda **kwargs: calls["post_call"].append(kwargs) + return logging_obj, calls + + +def test_rust_true_serves_the_call_and_stamps_the_header(): + seen = _inject() + response = _run() + + assert response.choices[0].message.content == "hello from rust" + assert response._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + assert len(seen["call"]) == 1 + + +def test_the_core_receives_the_credentials_this_handler_already_resolved(): + """Both paths must sign as the same principal, so the resolved credentials + are handed down rather than re-derived from ambient AWS state.""" + seen = _inject() + _run() + + params = seen["call"][0]["optional_params"] + assert params["aws_access_key_id"] == "AKIARESOLVED" + assert params["aws_secret_access_key"] == "resolved-secret" + assert params["aws_session_token"] == "resolved-token" + assert params["aws_region_name"] == "us-east-1" + + +def test_the_core_receives_the_converse_url_this_handler_already_built(): + seen = _inject() + _run() + + assert seen["call"][0]["api_base"].endswith( + "/model/anthropic.claude-sonnet-4-5-v1%3A0/converse" + ) + assert "bedrock-runtime.us-east-1.amazonaws.com" in seen["call"][0]["api_base"] + + +def test_the_core_receives_the_untranslated_openai_messages(): + seen = _inject() + _run( + messages=[ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + ] + ) + assert seen["call"][0]["messages"] == [ + {"role": "system", "content": "be terse"}, + {"role": "user", "content": "hi"}, + ] + + +def test_without_the_opt_in_the_core_is_never_consulted(): + seen = _inject() + try: + _run(litellm_params={}) + except Exception: + # The Python path goes on to make an HTTP call; not reaching the gate + # is the assertion, so a failure past this point is expected. + pass + assert seen["gate"] == [] + assert seen["call"] == [] + + +def test_streaming_stays_on_the_python_path(): + seen = _inject() + try: + _run(optional_params={"maxTokens": 16, "stream": True}) + except Exception: + pass + assert seen["gate"] == [] + + +def test_a_declined_request_never_reaches_the_native_call(): + seen = _inject(decline_reason="unrecognized request parameter") + try: + _run() + except Exception: + pass + assert len(seen["gate"]) == 1 + assert seen["call"] == [] + + +def test_pre_call_logging_fires_exactly_once_on_the_rust_path(): + _inject() + logging_obj = MagicMock() + _run(logging_obj=logging_obj) + assert logging_obj.pre_call.call_count == 1 + + +@pytest.mark.asyncio +async def test_the_async_path_falls_back_when_the_core_declines(monkeypatch): + class _Declined(Exception): + pass + + class _FakeNative: + RustBridgeDeclined = _Declined + RustUpstreamError = type("_Upstream", (Exception,), {}) + + monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + + async def declining_native(**_kwargs): + raise _Declined("blank message text") + + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=declining_native + ) + + sentinel = object() + + async def python_path(**_kwargs): + return sentinel + + with ( + patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ), + patch.object( + BedrockConverseLLM, "async_completion", side_effect=python_path + ) as python_call, + ): + result = await BedrockConverseLLM().completion( + **_completion_kwargs(acompletion=True) + ) + + assert result is sentinel + assert python_call.called, "a failing rust call must re-enter the python path" + + +@pytest.mark.asyncio +async def test_the_async_path_serves_the_rust_response_without_the_fallback(): + async def native(**_kwargs): + return dict(RUST_RESPONSE) + + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=native + ) + + with ( + patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ), + patch.object(BedrockConverseLLM, "async_completion") as python_call, + ): + result = await BedrockConverseLLM().completion( + **_completion_kwargs(acompletion=True) + ) + + assert result.choices[0].message.content == "hello from rust" + assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + assert not python_call.called + + +@pytest.mark.asyncio +async def test_pre_call_logging_fires_once_even_when_the_rust_path_declines(): + """One request, one pre_call. Without the suppression the Python fallback + logs a second one and non-idempotent callbacks run twice.""" + + class _Declined(Exception): + pass + + class _FakeNative: + RustBridgeDeclined = _Declined + RustUpstreamError = type("_Upstream", (Exception,), {}) + + async def declining_native(**_kwargs): + raise _Declined("blank message text") + + logging_obj = MagicMock() + served = [] + + async def python_path(**kwargs): + served.append(kwargs) + return ModelResponse() + + with ( + patch.object(bridge, "get_native_bridge", lambda: _FakeNative()), + patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ), + patch.object( + BedrockConverseLLM, "async_completion", side_effect=python_path + ), + ): + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=declining_native + ) + await BedrockConverseLLM().completion( + **_completion_kwargs(acompletion=True, logging_obj=logging_obj) + ) + + assert logging_obj.pre_call.call_count == 1 + assert served and served[0]["skip_pre_call_logging"] is True + + +CONVERSE_RESPONSE = { + "output": {"message": {"role": "assistant", "content": [{"text": "hi"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 5, "outputTokens": 2, "totalTokens": 7}, +} + + +async def _drive_async_completion(*, skip_pre_call_logging: bool, logging_obj): + """Run the real `async_completion` with a stubbed transport.""" + import httpx as _httpx + + client = MagicMock() + + async def post(**_kwargs): + return _httpx.Response( + 200, + json=CONVERSE_RESPONSE, + request=_httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com"), + ) + + client.post = post + client.__class__ = AsyncHTTPHandler + + return await BedrockConverseLLM().async_completion( + model="anthropic.claude-sonnet-4-5-v1:0", + messages=[{"role": "user", "content": "hi"}], + api_base="https://bedrock-runtime.us-west-2.amazonaws.com/model/m/converse", + model_response=ModelResponse(), + timeout=30.0, + encoding=None, + logging_obj=logging_obj, + stream=None, + optional_params={"maxTokens": 16}, + litellm_params={"aws_region_name": "us-west-2"}, + credentials=RESOLVED_CREDENTIALS, + headers={}, + client=client, + skip_pre_call_logging=skip_pre_call_logging, + ) + + +@pytest.mark.asyncio +async def test_async_completion_honors_the_pre_call_suppression(): + logging_obj = MagicMock() + await _drive_async_completion(skip_pre_call_logging=True, logging_obj=logging_obj) + assert logging_obj.pre_call.call_count == 0 + + +@pytest.mark.asyncio +async def test_async_completion_logs_pre_call_by_default(): + """The suppression must be opt-in, so every existing caller keeps its log.""" + logging_obj = MagicMock() + await _drive_async_completion(skip_pre_call_logging=False, logging_obj=logging_obj) + assert logging_obj.pre_call.call_count == 1 + + +def _sync_client_returning_converse_response(): + client = MagicMock() + client.post = lambda **_kwargs: httpx.Response( + 200, + json=CONVERSE_RESPONSE, + request=httpx.Request("POST", "https://bedrock-runtime.us-west-2.amazonaws.com"), + ) + client.__class__ = HTTPHandler + return client + + +def test_pre_call_logging_fires_once_when_the_sync_rust_path_declines(): + """One request, one pre_call, on the synchronous path too. + + The gate accepts and logs, then the native call declines before the + provider is reached, so execution continues into the Python path below. + That is the same attempt continuing; without the suppression it logs a + second pre_call and non-idempotent callbacks run twice for one request. + """ + + class _Declined(Exception): + pass + + class _FakeNative: + RustBridgeDeclined = _Declined + RustUpstreamError = type("_Upstream", (Exception,), {}) + + def declining_native(**_kwargs): + raise _Declined("blank message text") + + logging_obj = MagicMock() + + with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, chat_completions=declining_native + ) + response = _run( + logging_obj=logging_obj, + client=_sync_client_returning_converse_response(), + ) + + assert response.choices[0].message.content == "hi" + assert logging_obj.pre_call.call_count == 1 + + +def test_the_sync_python_path_still_logs_pre_call_without_the_opt_in(): + """The suppression must not swallow the log on a request the gate declined, + so a deployment with no `rust` flag keeps exactly the log it always had.""" + logging_obj = MagicMock() + response = _run( + logging_obj=logging_obj, + litellm_params={}, + client=_sync_client_returning_converse_response(), + ) + + assert response.choices[0].message.content == "hi" + assert logging_obj.pre_call.call_count == 1 + + +def test_post_call_logging_fires_on_the_sync_rust_path(): + """The Rust core owns the provider call, so the Converse transform that + normally raises `post_call` never runs. Without the bridge hook every + post_call callback goes silent and `original_response` stays unset.""" + import json + + _inject() + logging_obj = MagicMock() + _run(logging_obj=logging_obj) + + assert logging_obj.post_call.call_count == 1 + logged = logging_obj.post_call.call_args.kwargs["original_response"] + assert json.loads(logged)["choices"][0]["message"]["content"] == "hello from rust" + + +@pytest.mark.asyncio +async def test_post_call_logging_fires_on_the_async_rust_path(): + """The asynchronous path runs through the same hook, so the two paths + cannot drift apart the way the pre_call suppression once did.""" + import json + + async def native(**_kwargs): + return dict(RUST_RESPONSE) + + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, achat_completions=native + ) + logging_obj = MagicMock() + + with patch.object( + BedrockConverseLLM, "get_credentials", return_value=RESOLVED_CREDENTIALS + ): + await BedrockConverseLLM().completion( + **_completion_kwargs(acompletion=True, logging_obj=logging_obj) + ) + + assert logging_obj.post_call.call_count == 1 + logged = logging_obj.post_call.call_args.kwargs["original_response"] + assert json.loads(logged)["choices"][0]["message"]["content"] == "hello from rust" + + +def test_post_call_is_not_logged_twice_when_the_sync_rust_call_declines(): + """A decline never reached the provider, so the Python path serves the + request and owns the only post_call. Firing the hook there too would double + every post_call callback for one request.""" + + class _Declined(Exception): + pass + + class _FakeNative: + RustBridgeDeclined = _Declined + RustUpstreamError = type("_Upstream", (Exception,), {}) + + def declining_native(**_kwargs): + raise _Declined("blank message text") + + logging_obj, calls = _recording_logging_obj() + + with patch.object(bridge, "get_native_bridge", lambda: _FakeNative()): + bridge.set_rust_chat_completions( + decline=lambda **_kwargs: None, chat_completions=declining_native + ) + response = _run( + logging_obj=logging_obj, + client=_sync_client_returning_converse_response(), + ) + + assert response.choices[0].message.content == "hi" + assert len(calls["post_call"]) == 1 + assert "hi" in calls["post_call"][0]["original_response"] diff --git a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py index a3be3ebcfc7..ebc482a44ba 100644 --- a/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py +++ b/tests/test_litellm/llms/bedrock/chat/test_converse_transformation.py @@ -3011,7 +3011,7 @@ def test_request_metadata_validation(): litellm_params={}, headers={}, ) - assert False, "Should have raised validation error for too many items" + pytest.fail("Should have raised validation error for too many items") except Exception as e: assert "maximum of 16 items" in str(e).lower() @@ -3034,7 +3034,7 @@ def test_request_metadata_key_constraints(): litellm_params={}, headers={}, ) - assert False, "Should have raised validation error for key too long" + pytest.fail("Should have raised validation error for key too long") except Exception as e: assert "key length" in str(e).lower() or "256 characters" in str(e).lower() @@ -3049,7 +3049,7 @@ def test_request_metadata_key_constraints(): litellm_params={}, headers={}, ) - assert False, "Should have raised validation error for empty key" + pytest.fail("Should have raised validation error for empty key") except Exception as e: assert "key length" in str(e).lower() or "empty" in str(e).lower() @@ -3072,7 +3072,7 @@ def test_request_metadata_value_constraints(): litellm_params={}, headers={}, ) - assert False, "Should have raised validation error for value too long" + pytest.fail("Should have raised validation error for value too long") except Exception as e: assert "value length" in str(e).lower() or "256 characters" in str(e).lower() diff --git a/tests/test_litellm/llms/custom_httpx/test_container_handler.py b/tests/test_litellm/llms/custom_httpx/test_container_handler.py new file mode 100644 index 00000000000..a1b5a66696d --- /dev/null +++ b/tests/test_litellm/llms/custom_httpx/test_container_handler.py @@ -0,0 +1,102 @@ +from unittest.mock import MagicMock + +import httpx +import pytest + +import litellm +from litellm.llms.base_llm.chat.transformation import BaseLLMException +from litellm.llms.custom_httpx.container_handler import generic_container_handler +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import ProviderConfigManager + +FILE_NOT_FOUND_BODY = { + "error": { + "message": "File not found.", + "type": "invalid_request_error", + "param": None, + "code": None, + } +} + + +def _sync_client(response: httpx.Response) -> HTTPHandler: + handler = HTTPHandler() + handler.client = httpx.Client(transport=httpx.MockTransport(lambda _request: response)) + return handler + + +def _async_client(response: httpx.Response) -> AsyncHTTPHandler: + handler = AsyncHTTPHandler() + handler.client = httpx.AsyncClient(transport=httpx.MockTransport(lambda _request: response)) + return handler + + +def _handle(endpoint_name: str, client, **overrides): + return generic_container_handler.handle( + endpoint_name=endpoint_name, + container_provider_config=ProviderConfigManager.get_provider_container_config( + provider=litellm.LlmProviders.OPENAI + ), + litellm_params=GenericLiteLLMParams(api_key="sk-test"), + logging_obj=MagicMock(), + client=client, + container_id="cntr_real", + file_id="cfile_nonexistent", + **overrides, + ) + + +def test_binary_endpoint_raises_on_error_status(): + with pytest.raises(BaseLLMException) as exc_info: + _handle( + "retrieve_container_file_content", + _sync_client(httpx.Response(404, json=FILE_NOT_FOUND_BODY)), + ) + + assert exc_info.value.status_code == 404 + assert exc_info.value.message == "File not found." + + +@pytest.mark.asyncio +async def test_async_binary_endpoint_raises_on_error_status(): + with pytest.raises(BaseLLMException) as exc_info: + await _handle( + "aretrieve_container_file_content", + _async_client(httpx.Response(404, json=FILE_NOT_FOUND_BODY)), + _is_async=True, + ) + + assert exc_info.value.status_code == 404 + assert exc_info.value.message == "File not found." + + +def test_binary_endpoint_returns_raw_content_on_success(): + content = _handle( + "retrieve_container_file_content", + _sync_client(httpx.Response(200, content=b"\x00binary-payload")), + ) + + assert content == b"\x00binary-payload" + + +def test_error_status_with_non_json_body_surfaces_response_text(): + with pytest.raises(BaseLLMException) as exc_info: + _handle( + "retrieve_container_file_content", + _sync_client(httpx.Response(502, content=b"bad gateway")), + ) + + assert exc_info.value.status_code == 502 + assert exc_info.value.message == "bad gateway" + + +def test_json_endpoint_still_raises_provider_error_message(): + with pytest.raises(BaseLLMException) as exc_info: + _handle( + "retrieve_container_file", + _sync_client(httpx.Response(404, json=FILE_NOT_FOUND_BODY)), + ) + + assert exc_info.value.status_code == 404 + assert exc_info.value.message == "File not found." diff --git a/tests/test_litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py b/tests/test_litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py index d5783e3567f..fa6f23dc7ff 100644 --- a/tests/test_litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py +++ b/tests/test_litellm/llms/deepseek/chat/test_deepseek_chat_transformation.py @@ -106,3 +106,184 @@ async def test_async_transform_request_strips_unsupported_tools_from_body(): def test_thinking_mode_active_bool_thinking_returns_false_without_crashing(): config = DeepSeekChatConfig() assert config._thinking_mode_active(model="deepseek-reasoner", optional_params={"thinking": True}) is False + + +class TestDeepSeekThinkingParams: + """Test thinking and reasoning_effort parameter handling for DeepSeek.""" + + def setup_method(self): + self.config = DeepSeekChatConfig() + self.model = "deepseek-reasoner" + + def test_get_supported_openai_params_includes_thinking(self): + """Test that thinking and reasoning_effort are in supported params.""" + params = self.config.get_supported_openai_params(self.model) + assert "thinking" in params + assert "reasoning_effort" in params + + def test_map_thinking_enabled(self): + """Test that thinking={"type": "enabled"} is passed through correctly.""" + non_default_params = {"thinking": {"type": "enabled"}} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert result["thinking"] == {"type": "enabled"} + + def test_map_thinking_with_budget_tokens_strips_budget(self): + """Test that budget_tokens is stripped from thinking param (DeepSeek doesn't support it).""" + non_default_params = {"thinking": {"type": "enabled", "budget_tokens": 2048}} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + # Should strip budget_tokens, only pass type + assert result["thinking"] == {"type": "enabled"} + assert "budget_tokens" not in result.get("thinking", {}) + + def test_map_reasoning_effort_medium(self): + """Test that reasoning_effort='medium' maps to thinking enabled.""" + non_default_params = {"reasoning_effort": "medium"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert result["thinking"] == {"type": "enabled"} + + def test_map_reasoning_effort_low(self): + """Test that reasoning_effort='low' maps to thinking enabled.""" + non_default_params = {"reasoning_effort": "low"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert result["thinking"] == {"type": "enabled"} + + def test_map_reasoning_effort_high(self): + """Test that reasoning_effort='high' maps to thinking enabled.""" + non_default_params = {"reasoning_effort": "high"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert result["thinking"] == {"type": "enabled"} + + def test_map_reasoning_effort_none_does_not_enable_thinking(self): + """Test that reasoning_effort='none' does not enable thinking.""" + non_default_params = {"reasoning_effort": "none"} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert result["thinking"] == {"type": "disabled"} + + def test_map_reasoning_effort_null_does_not_enable_thinking(self): + """Test that reasoning_effort=None does not enable thinking.""" + non_default_params = {"reasoning_effort": None} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert "thinking" not in result + + def test_thinking_takes_precedence_over_reasoning_effort(self): + """Test that thinking param takes precedence when both are provided.""" + non_default_params = { + "thinking": {"type": "enabled"}, + "reasoning_effort": "high", + } + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + # thinking should be set, reasoning_effort should not override + assert result["thinking"] == {"type": "enabled"} + + def test_invalid_thinking_type_ignored(self): + """Test that invalid thinking type values are ignored.""" + non_default_params = {"thinking": {"type": "invalid"}} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert "thinking" not in result + + def test_thinking_none_value_ignored(self): + """Test that thinking=None is ignored.""" + non_default_params = {"thinking": None} + optional_params = {} + + result = self.config.map_openai_params( + non_default_params=non_default_params, + optional_params=optional_params, + model=self.model, + drop_params=False, + ) + + assert "thinking" not in result + + def test_drop_unsupported_tools_removes_dangling_tool_choice(self): + optional_params = { + "tools": [ + {"type": "namespace", "name": "local_shell"}, + {"type": "function", "function": {"name": "get_weather"}}, + ], + "tool_choice": { + "type": "function", + "function": {"name": "local_shell"}, + }, + "parallel_tool_calls": True, + } + + result = self.config._drop_unsupported_tools(optional_params) + + assert result["tools"] == [ + {"type": "function", "function": {"name": "get_weather"}} + ] + assert "tool_choice" not in result + assert result["parallel_tool_calls"] is True diff --git a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py index 53c9e4b207c..0f0033cae36 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py @@ -47,6 +47,30 @@ def supplied_params(request): return request.param + +_AMBIENT_OCI_ENV: tuple[str, ...] = ( + "OCI_REGION", + "OCI_USER", + "OCI_FINGERPRINT", + "OCI_TENANCY", + "OCI_KEY_FILE", + "OCI_KEY", + "OCI_COMPARTMENT_ID", +) + + +@pytest.fixture +def without_ambient_oci_env(monkeypatch): + """Drop OCI credentials the environment may supply. + + validate_environment falls back to os.environ for every credential and to a + default region only when OCI_REGION is unset, so a developer or runner with + OCI configured would see these tests find credentials they never passed. + """ + for variable in _AMBIENT_OCI_ENV: + monkeypatch.delenv(variable, raising=False) + +@pytest.mark.usefixtures("without_ambient_oci_env") class TestOCIChatConfig: def test_validate_environment_with_oci_region(self, supplied_params): config = OCIChatConfig() @@ -1552,3 +1576,330 @@ class TestOCIChatConfigErrorPaths: import pytest from unittest.mock import MagicMock +from litellm.llms.oci.common_utils import OCIError, sign_with_manual_credentials + + + +@pytest.fixture +def config(): + return OCIChatConfig() + + +@pytest.mark.usefixtures("without_ambient_oci_env") +class TestOCIKeyNormalization: + """Tests for OCI private key content normalization.""" + + def test_oci_key_with_escaped_newlines(self, config): + """Test that escaped newlines (\\n) are converted to actual newlines.""" + # Simulate PEM content with escaped newlines (as would come from JSON/UI input) + escaped_pem = "-----BEGIN RSA PRIVATE KEY-----\\nMIIEowIBAAKCAQEA...\\n-----END RSA PRIVATE KEY-----" + + optional_params = { + "oci_user": "ocid1.user.oc1..test", + "oci_fingerprint": "aa:bb:cc:dd", + "oci_tenancy": "ocid1.tenancy.oc1..test", + "oci_region": "us-ashburn-1", + "oci_key": escaped_pem, + } + + # We can't fully test signing without a real key, but we can verify + # the error message indicates the key was processed (not a type error) + with pytest.raises(Exception) as exc_info: + sign_with_manual_credentials( + headers={}, + optional_params=optional_params, + request_data={"test": "data"}, + api_base="https://test.oci.oraclecloud.com/api", + ) + + # The error should be about key format/loading, not about type + # This confirms the string was processed and newlines were normalized + error_message = str(exc_info.value) + assert "must be a string" not in error_message.lower() + + def test_oci_key_with_crlf_newlines(self, config): + """Test that Windows-style CRLF newlines are normalized to LF.""" + # Simulate PEM content with CRLF newlines + crlf_pem = "-----BEGIN RSA PRIVATE KEY-----\r\nMIIEowIBAAKCAQEA...\r\n-----END RSA PRIVATE KEY-----" + + optional_params = { + "oci_user": "ocid1.user.oc1..test", + "oci_fingerprint": "aa:bb:cc:dd", + "oci_tenancy": "ocid1.tenancy.oc1..test", + "oci_region": "us-ashburn-1", + "oci_key": crlf_pem, + } + + with pytest.raises(Exception) as exc_info: + sign_with_manual_credentials( + headers={}, + optional_params=optional_params, + request_data={"test": "data"}, + api_base="https://test.oci.oraclecloud.com/api", + ) + + error_message = str(exc_info.value) + assert "must be a string" not in error_message.lower() + + def test_oci_key_rejects_non_string_type(self, config): + """Test that non-string oci_key values raise OCIError.""" + optional_params = { + "oci_user": "ocid1.user.oc1..test", + "oci_fingerprint": "aa:bb:cc:dd", + "oci_tenancy": "ocid1.tenancy.oc1..test", + "oci_region": "us-ashburn-1", + "oci_key": {"invalid": "dict"}, # Wrong type + } + + with pytest.raises(OCIError) as exc_info: + sign_with_manual_credentials( + headers={}, + optional_params=optional_params, + request_data={"test": "data"}, + api_base="https://test.oci.oraclecloud.com/api", + ) + + assert exc_info.value.status_code == 400 + assert "must be a string" in str(exc_info.value.message) + assert "dict" in str(exc_info.value.message) + + def test_oci_key_rejects_list_type(self, config): + """Test that list oci_key values raise OCIError.""" + optional_params = { + "oci_user": "ocid1.user.oc1..test", + "oci_fingerprint": "aa:bb:cc:dd", + "oci_tenancy": "ocid1.tenancy.oc1..test", + "oci_region": "us-ashburn-1", + "oci_key": ["invalid", "list"], # Wrong type + } + + with pytest.raises(OCIError) as exc_info: + sign_with_manual_credentials( + headers={}, + optional_params=optional_params, + request_data={"test": "data"}, + api_base="https://test.oci.oraclecloud.com/api", + ) + + assert exc_info.value.status_code == 400 + assert "must be a string" in str(exc_info.value.message) + assert "list" in str(exc_info.value.message) + + +@pytest.mark.usefixtures("without_ambient_oci_env") +class TestOCIValidateEnvironment: + """Tests for OCI environment validation.""" + + def test_missing_required_credentials_raises_error(self, config): + """Test that missing required credentials raise an error.""" + with pytest.raises(Exception) as exc_info: + config.validate_environment( + headers={}, + model="oci/xai.grok-3", + messages=[{"role": "user", "content": "Hello"}], + optional_params={}, # No credentials provided + litellm_params={}, + api_key=None, + api_base=None, + ) + + error_message = str(exc_info.value) + assert "oci_user" in error_message + assert "oci_fingerprint" in error_message + assert "oci_tenancy" in error_message + + def test_validate_environment_with_all_credentials(self, config): + """Test that validation passes with all required credentials.""" + headers = config.validate_environment( + headers={}, + model="oci/xai.grok-3", + messages=[{"role": "user", "content": "Hello"}], + optional_params={ + "oci_user": "ocid1.user.oc1..test", + "oci_fingerprint": "aa:bb:cc:dd", + "oci_tenancy": "ocid1.tenancy.oc1..test", + "oci_region": "us-ashburn-1", + "oci_compartment_id": "ocid1.compartment.oc1..test", + "oci_key": "-----BEGIN RSA PRIVATE KEY-----\ntest\n-----END RSA PRIVATE KEY-----", + }, + litellm_params={}, + api_key=None, + api_base=None, + ) + + assert headers["content-type"] == "application/json" + assert "user-agent" in headers + + +@pytest.mark.usefixtures("without_ambient_oci_env") +class TestOCIGetCompleteUrl: + """Tests for OCI URL generation.""" + + def test_get_complete_url_default_region(self, config): + """Test URL generation with default region.""" + url = config.get_complete_url( + api_base=None, + api_key=None, + model="oci/xai.grok-3", + optional_params={}, + litellm_params={}, + stream=False, + ) + + assert "us-ashburn-1" in url + assert "inference.generativeai" in url + assert "/20231130/actions/chat" in url + + def test_get_complete_url_custom_region(self, config): + """Test URL generation with custom region.""" + url = config.get_complete_url( + api_base=None, + api_key=None, + model="oci/xai.grok-3", + optional_params={"oci_region": "eu-frankfurt-1"}, + litellm_params={}, + stream=False, + ) + + assert "eu-frankfurt-1" in url + assert "inference.generativeai" in url + + +@pytest.mark.usefixtures("without_ambient_oci_env") +class TestOCIImageUrlTransformation: + """Tests for OCI image_url format handling in multimodal messages. + + Fixes: https://github.com/BerriAI/litellm/issues/18270 + Fixes: https://github.com/BerriAI/litellm/issues/19589 + """ + + def test_image_url_as_string(self): + """Test that image_url as a plain string works.""" + from litellm.llms.oci.chat.transformation import ( + adapt_messages_to_generic_oci_standard, + ) + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this image?"}, + {"type": "image_url", "image_url": "https://example.com/image.png"}, + ], + } + ] + + result = adapt_messages_to_generic_oci_standard(messages) + + assert len(result) == 1 + assert result[0].role == "USER" + assert len(result[0].content) == 2 + # imageUrl is now an OCIImageUrl object with a 'url' property + assert result[0].content[1].imageUrl.url == "https://example.com/image.png" + + def test_image_url_as_openai_object(self): + """Test that image_url as OpenAI-style object {"url": "..."} works.""" + from litellm.llms.oci.chat.transformation import ( + adapt_messages_to_generic_oci_standard, + ) + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this image?"}, + { + "type": "image_url", + "image_url": {"url": "https://example.com/image.png"}, + }, + ], + } + ] + + result = adapt_messages_to_generic_oci_standard(messages) + + assert len(result) == 1 + assert result[0].role == "USER" + assert len(result[0].content) == 2 + # imageUrl is now an OCIImageUrl object with a 'url' property + assert result[0].content[1].imageUrl.url == "https://example.com/image.png" + + def test_image_url_serializes_as_object(self): + """Test that imageUrl serializes as {"url": "..."} for OCI API. + + Fixes: https://github.com/BerriAI/litellm/issues/19589 + OCI expects imageUrl to be an object with a 'url' property, not a plain string. + """ + from litellm.llms.oci.chat.transformation import ( + adapt_messages_to_generic_oci_standard, + ) + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "Describe this image."}, + { + "type": "image_url", + "image_url": {"url": "data:image/png;base64,ABC123"}, + }, + ], + } + ] + + result = adapt_messages_to_generic_oci_standard(messages) + image_part = result[0].content[1] + + # Serialize as OCI would receive it (with exclude_none=True) + serialized = image_part.model_dump(exclude_none=True) + + # Verify the structure matches OCI's expected format + assert serialized == { + "type": "IMAGE", + "imageUrl": {"url": "data:image/png;base64,ABC123"}, + } + + def test_image_url_invalid_type_raises_error(self): + """Test that invalid image_url type raises an error.""" + from litellm.llms.oci.chat.transformation import ( + adapt_messages_to_generic_oci_standard, + ) + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this image?"}, + {"type": "image_url", "image_url": 12345}, # Invalid type + ], + } + ] + + with pytest.raises(Exception) as exc_info: + adapt_messages_to_generic_oci_standard(messages) + + assert "image_url" in str(exc_info.value) + + def test_image_url_object_missing_url_raises_error(self): + """Test that object without 'url' property raises an error.""" + from litellm.llms.oci.chat.transformation import ( + adapt_messages_to_generic_oci_standard, + ) + + messages = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "What is in this image?"}, + { + "type": "image_url", + "image_url": {"detail": "high"}, + }, # Missing 'url' + ], + } + ] + + with pytest.raises(Exception) as exc_info: + adapt_messages_to_generic_oci_standard(messages) + + assert "image_url" in str(exc_info.value) diff --git a/tests/test_litellm/llms/oci/test_oci_coverage_boost.py b/tests/test_litellm/llms/oci/test_oci_coverage_boost.py index 0b7afa3775d..7c91ece70b5 100644 --- a/tests/test_litellm/llms/oci/test_oci_coverage_boost.py +++ b/tests/test_litellm/llms/oci/test_oci_coverage_boost.py @@ -11,11 +11,16 @@ All tests are self-contained and require no real OCI credentials or network acce """ import json +from typing import TYPE_CHECKING + import pytest from unittest.mock import patch, MagicMock, AsyncMock import httpx +if TYPE_CHECKING: + from litellm.llms.oci.chat.transformation import OCIStreamWrapper + from litellm import ModelResponse from litellm.llms.oci.chat.cohere import ( _extract_text_content, diff --git a/tests/test_litellm/llms/openai/responses/test_openai_count_tokens_transformation.py b/tests/test_litellm/llms/openai/responses/test_openai_count_tokens_transformation.py index a87aaa7435f..bcca35886fe 100644 --- a/tests/test_litellm/llms/openai/responses/test_openai_count_tokens_transformation.py +++ b/tests/test_litellm/llms/openai/responses/test_openai_count_tokens_transformation.py @@ -1,6 +1,8 @@ import os import sys +import pytest + sys.path.insert( 0, os.path.abspath("../../../../..") ) # Adds the parent directory to the system path @@ -177,7 +179,7 @@ def test_validate_request_missing_model(): config = OpenAICountTokensConfig() try: config.validate_request(model="", input="Hello") - assert False, "Should have raised ValueError" + pytest.fail("Should have raised ValueError") except ValueError as e: assert "model" in str(e) @@ -187,7 +189,7 @@ def test_validate_request_missing_input(): config = OpenAICountTokensConfig() try: config.validate_request(model="gpt-4o", input="") - assert False, "Should have raised ValueError" + pytest.fail("Should have raised ValueError") except ValueError as e: assert "input" in str(e) diff --git a/tests/test_litellm/llms/openrouter/responses/test_openrouter_responses_transformation.py b/tests/test_litellm/llms/openrouter/responses/test_openrouter_responses_transformation.py index 62c372a3d93..d279eaeec01 100644 --- a/tests/test_litellm/llms/openrouter/responses/test_openrouter_responses_transformation.py +++ b/tests/test_litellm/llms/openrouter/responses/test_openrouter_responses_transformation.py @@ -9,6 +9,8 @@ reasoning.encrypted_content for multi-turn stateless workflows. Related issue: https://github.com/BerriAI/litellm/issues/22189 """ +import pytest + import litellm from litellm.llms.openrouter.responses.transformation import ( OpenRouterResponsesAPIConfig, @@ -76,7 +78,7 @@ class TestOpenRouterResponsesAPIConfig: model="openai/o4-mini", litellm_params=GenericLiteLLMParams(), ) - assert False, "Should have raised ValueError" + pytest.fail("Should have raised ValueError") except ValueError as e: assert "OpenRouter API key is required" in str(e) diff --git a/tests/test_litellm/llms/perplexity/embedding/test_perplexity_embedding_transformation.py b/tests/test_litellm/llms/perplexity/embedding/test_perplexity_embedding_transformation.py index 15ebecdcb1d..89c0ec1988f 100644 --- a/tests/test_litellm/llms/perplexity/embedding/test_perplexity_embedding_transformation.py +++ b/tests/test_litellm/llms/perplexity/embedding/test_perplexity_embedding_transformation.py @@ -8,6 +8,7 @@ import struct from unittest.mock import MagicMock import httpx +import pytest from litellm.llms.perplexity.embedding.transformation import ( PerplexityEmbeddingConfig, @@ -245,7 +246,7 @@ class TestPerplexityEmbeddingConfig: model_response=model_response, logging_obj=self.logging_obj, ) - assert False, "Should have raised PerplexityEmbeddingError" + pytest.fail("Should have raised PerplexityEmbeddingError") except PerplexityEmbeddingError as e: assert e.status_code == 500 assert "Server error" in e.message diff --git a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py index 46c1e457d7c..16708e062e4 100644 --- a/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py +++ b/tests/test_litellm/llms/perplexity/test_perplexity_cost_calculator.py @@ -400,6 +400,31 @@ class TestPerplexityCostCalculator: assert completion_cost == 0.008 assert prompt_cost + completion_cost == 0.008 + def test_uses_perplexity_provided_cost_when_normalized_to_float(self): + """ + Regression: for Responses API / Agent API models, `ResponseAPIUsage.parse_cost` + (litellm/types/llms/openai.py) already flattens Perplexity's + `usage.cost.total_cost` dict down to a plain float before + `_transform_response_api_usage_to_chat_usage` (litellm/responses/utils.py) copies + it onto the chat `Usage` object. So `usage.cost` arrives here as a float, not a + dict, on that path. + + Pre-fix, the `isinstance(cost_info, dict)` check was always False for a float, + so the pre-calculated cost branch was dead code for every Responses-mode + Perplexity model and it silently fell back to manual token-rate calculation, + recording $0 for any model missing static per-token rates (e.g. + perplexity/openai/gpt-5.2 before rates existed). + """ + usage = Usage(prompt_tokens=100, completion_tokens=50, total_tokens=150) + usage.cost = 0.008 + + prompt_cost, completion_cost = perplexity_cost_per_token( + model="sonar-pro", usage=usage + ) + + assert prompt_cost == 0.0 + assert completion_cost == 0.008 + def test_falls_back_to_manual_calculation_when_no_cost_provided(self): """ Test that manual cost calculation is used when Perplexity doesn't @@ -451,3 +476,52 @@ class TestPerplexityCostCalculator: assert math.isclose(prompt_cost, expected_prompt, rel_tol=1e-9) assert math.isclose(completion_cost, expected_completion, rel_tol=1e-9) + + @pytest.mark.parametrize( + "model_id, usd_per_1m_input, usd_per_1m_output, usd_per_1m_cache_read", + [ + ("deepseek-v4-flash-0731", 0.13, 0.26, 0.028), + ("glm-5.2", 1.4, 4.4, 0.14), + ("kimi-k3", 3.0, 15.0, 0.3), + ("kimi-k2.7-code", 0.95, 4.0, 0.19), + ], + ) + def test_agent_api_entries_carry_perplexity_published_rates( + self, model_id, usd_per_1m_input, usd_per_1m_output, usd_per_1m_cache_read + ): + """The Agent API third-party models are priced from Perplexity's own catalog + (GET https://api.perplexity.ai/v1/models, `pricing` in usd_per_1m_tokens). + Perplexity's model id already starts with `perplexity/`, so the cost-map key + doubles the prefix. Regression: glm-5.2 shipped glm-5.3's 0.26 cache-read rate, + copied from the neighbouring catalog row, an 86% overcharge on cached input. + """ + info = get_model_info( + model=f"perplexity/{model_id}", custom_llm_provider="perplexity" + ) + + assert info["key"] == f"perplexity/perplexity/{model_id}" + assert info["litellm_provider"] == "perplexity" + assert info["mode"] == "responses" + assert math.isclose(info["input_cost_per_token"], usd_per_1m_input / 1e6, rel_tol=1e-9) + assert math.isclose(info["output_cost_per_token"], usd_per_1m_output / 1e6, rel_tol=1e-9) + assert math.isclose( + info["cache_read_input_token_cost"], usd_per_1m_cache_read / 1e6, rel_tol=1e-9 + ) + + def test_agent_api_fallback_rates_price_a_response_without_metered_cost(self): + """Perplexity meters cost on the response, but when `usage.cost` is absent the + calculator falls back to the mapped per-token rates. Regression: that fallback + raised "This model isn't mapped yet" for every Agent API third-party model, + because the doubled cost-map key was unreachable from the resolution ladder. + """ + from litellm import ModelResponse + + response = ModelResponse() + response.usage = Usage(prompt_tokens=1000, completion_tokens=500, total_tokens=1500) + response.model = "perplexity/perplexity/glm-5.2" + + total_cost = completion_cost( + completion_response=response, custom_llm_provider="perplexity" + ) + + assert math.isclose(total_cost, 1000 * 1.4e-06 + 500 * 4.4e-06, rel_tol=1e-9) diff --git a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py index 7e13459bca1..f928964dab8 100644 --- a/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py +++ b/tests/test_litellm/llms/sagemaker/test_sagemaker_common_utils.py @@ -87,7 +87,6 @@ def test_sagemaker_response_stream_shape_is_structure_shape(): assert ( shape is not None ), "get_sagemaker_response_stream_shape() is None — botocore may not be installed" - shape: StructureShape = shape # remove Optional assert isinstance(shape, StructureShape) assert shape.name == "InvokeEndpointWithResponseStreamOutput" diff --git a/tests/test_litellm/passthrough/test_passthrough_main.py b/tests/test_litellm/passthrough/test_passthrough_main.py index 0b5bfac87bb..58a0185ea8c 100644 --- a/tests/test_litellm/passthrough/test_passthrough_main.py +++ b/tests/test_litellm/passthrough/test_passthrough_main.py @@ -43,9 +43,10 @@ def test_llm_passthrough_route(): client=client, ) - mock_post.call_args.kwargs[ - "request" - ].url == "http://localhost:8090/v1/chat/completions" + assert ( + mock_post.call_args.kwargs["request"].url + == "http://localhost:8090/v1/chat/completions" + ) assert response.status_code == 200 assert response.json == {"message": "Hello, world!"} diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py index 64afa52ab55..65e2faee1b2 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_list_outcomes.py @@ -2,6 +2,11 @@ to exactly one category, wire values never carry upstream prose, and single-upstream HTTP statuses stay truthful to who failed.""" +import sys + +if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11 + from exceptiongroup import BaseExceptionGroup + import httpx import pytest from mcp import McpError diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_traversal.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_traversal.py index a12c02339e6..b8bf4da1dc4 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_traversal.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_traversal.py @@ -2,6 +2,11 @@ links win (the ``raise ... from`` cause subtree, then ExceptionGroup members in raise order, then the incidental ``__context__`` chain last), and adversarial shapes terminate.""" +import sys + +if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11 + from exceptiongroup import BaseExceptionGroup + from litellm.proxy._experimental.mcp_server.faults import iter_exception_tree diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py index 5a2e65e7f68..50248e95ffa 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_db_credentials.py @@ -535,6 +535,164 @@ async def test_byok_guard_allows_overwriting_existing_oauth(): assert _stored_value(prisma) != oauth_row.credential_b64 +# ── Recovery from an unplanned LITELLM_SALT_KEY change ──────────────────────── + +PREVIOUS_SALT_KEY = "the-salt-key-this-deployment-used-before-9999" + + +def _row_written_under_previous_salt_key(monkeypatch, payload: str): + """A row encrypted under a salt key the proxy no longer holds. + + Asserts the fixture really is undecryptable under the current key, so a test + built on it cannot pass by accident. + """ + monkeypatch.setenv("LITELLM_SALT_KEY", PREVIOUS_SALT_KEY) + encrypted = encrypt_value_helper(payload) + monkeypatch.setenv("LITELLM_SALT_KEY", SALT_KEY) + assert _decode_user_credential(encrypted) is None, "fixture must not decrypt under the current salt key" + row = MagicMock() + row.credential_b64 = encrypted + row.user_id = "alice" + row.server_id = "srv-1" + return row + + +@pytest.mark.asyncio +async def test_reauthorization_replaces_row_written_under_previous_salt_key(monkeypatch): + # The wedged user: their row cannot be decrypted, so refusing preserves nothing. + old_payload = json.dumps({"type": "oauth2", "access_token": "tok-written-before-rotation"}) + prisma = _make_prisma_with_existing(row=_row_written_under_previous_salt_key(monkeypatch, old_payload)) + + await store_user_oauth_credential(prisma, "alice", "srv-1", "tok-after-reauthorization") + + # The replacement must decrypt under the CURRENT key and be the newly authorized token. + replacement = MagicMock() + replacement.credential_b64 = _stored_value(prisma) + replacement.server_id = "srv-1" + prisma.db.litellm_mcpusercredentials.find_unique = AsyncMock(return_value=replacement) + stored = await get_user_oauth_credential(prisma, "alice", "srv-1") + assert stored is not None + assert stored["access_token"] == "tok-after-reauthorization" + + +@pytest.mark.asyncio +async def test_readable_byok_is_still_refused_after_a_salt_key_change(monkeypatch): + # A legacy plain-base64 BYOK secret stays readable across a salt-key change, so + # the recovery path must not use it as an excuse to clobber a live credential. + monkeypatch.setenv("LITELLM_SALT_KEY", "a-completely-different-salt-key-4321") + prisma = _make_prisma_with_existing(row=_legacy_row("sk-live-byok-secret")) + + with pytest.raises(ValueError, match="could not be verified as an OAuth2"): + await store_user_oauth_credential(prisma, "alice", "srv-1", "tok") + + prisma.db.litellm_mcpusercredentials.upsert.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_recovery_warns_with_identifiers_and_never_logs_credentials(monkeypatch, caplog): + import logging + + old_payload = json.dumps({"type": "oauth2", "access_token": "tok-written-before-rotation"}) + row = _row_written_under_previous_salt_key(monkeypatch, old_payload) + prisma = _make_prisma_with_existing(row=row) + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + await store_user_oauth_credential(prisma, "alice", "srv-1", "tok-after-reauthorization") + + messages = [rec.getMessage() for rec in caplog.records] + matching = [m for m in messages if "could not be decrypted" in m and "replacing it" in m] + assert len(matching) == 1, f"expected one recovery warning, got {messages}" + assert "user=alice" in matching[0] and "server=srv-1" in matching[0] + for secret in ("tok-after-reauthorization", "tok-written-before-rotation", row.credential_b64): + assert secret not in matching[0] + + +@pytest.mark.asyncio +async def test_get_user_oauth_credential_warns_when_row_cannot_be_decrypted(monkeypatch, caplog): + import logging + + old_payload = json.dumps({"type": "oauth2", "access_token": "tok-written-before-rotation"}) + prisma = _make_prisma_with_existing(row=_row_written_under_previous_salt_key(monkeypatch, old_payload)) + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + assert await get_user_oauth_credential(prisma, "alice", "srv-1") is None + + matching = [rec.getMessage() for rec in caplog.records if "could not be decrypted" in rec.getMessage()] + assert len(matching) == 1, f"expected one read-path warning, got {[r.getMessage() for r in caplog.records]}" + assert "user=alice" in matching[0] and "server=srv-1" in matching[0] + + +@pytest.mark.asyncio +async def test_list_user_oauth_credentials_warns_per_row_when_rows_cannot_be_decrypted(monkeypatch, caplog): + # The bulk prefetch is the other read path, and it is by definition the multi-server case: + # a warning naming the wrong server sends the operator to the wrong place. Two wedged rows + # plus one healthy one, so a warning built from a constant or from the first row is caught. + import logging + + old_payload = json.dumps({"type": "oauth2", "access_token": "tok-written-before-rotation"}) + wedged_one = _row_written_under_previous_salt_key(monkeypatch, old_payload) + wedged_two = _row_written_under_previous_salt_key(monkeypatch, old_payload) + wedged_two.server_id = "srv-2" + + prisma = _make_prisma_with_existing(row=None) + await store_user_oauth_credential(prisma, "alice", "srv-3", "tok-healthy") + healthy = MagicMock() + healthy.credential_b64 = _stored_value(prisma) + healthy.server_id = "srv-3" + prisma.db.litellm_mcpusercredentials.find_many = AsyncMock(return_value=[wedged_one, healthy, wedged_two]) + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + result = await list_user_oauth_credentials(prisma, "alice") + + assert [cred["server_id"] for cred in result] == ["srv-3"] + matching = [rec.getMessage() for rec in caplog.records if "could not be decrypted" in rec.getMessage()] + assert len(matching) == 2, f"expected one warning per wedged row, got {matching}" + assert all("user=alice" in message for message in matching) + assert {"srv-1", "srv-2"} == {message.split("server=")[1].split(" ")[0] for message in matching} + + +@pytest.mark.asyncio +async def test_skip_byok_guard_does_not_read_the_existing_row(monkeypatch): + # The refresh paths pass skip_byok_guard=True precisely to save a DB round-trip on the + # hottest MCP path, so the flag has to actually suppress the lookup, not just the raise. + prisma = _make_prisma_with_existing(row=_legacy_row("plain-byok-key")) + + await store_user_oauth_credential(prisma, "alice", "srv-1", "tok", skip_byok_guard=True) + + prisma.db.litellm_mcpusercredentials.find_unique.assert_not_awaited() + prisma.db.litellm_mcpusercredentials.upsert.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_blank_credential_row_is_replaced_rather_than_refused(): + # A blank value decodes to "" rather than None, so it is not a decryption failure, but it + # holds no secret either. Pinned deliberately: the guard exists to protect readable + # content, and refusing here would wedge the user while preserving nothing. + blank = MagicMock() + blank.credential_b64 = "" + blank.user_id = "alice" + blank.server_id = "srv-1" + assert _decode_user_credential(blank.credential_b64) == "", "fixture must decode to empty, not None" + prisma = _make_prisma_with_existing(row=blank) + + await store_user_oauth_credential(prisma, "alice", "srv-1", "tok-after-reauthorization") + + prisma.db.litellm_mcpusercredentials.upsert.assert_awaited_once() + + +@pytest.mark.asyncio +async def test_readable_byok_row_does_not_warn_on_the_read_path(caplog): + # A BYOK row is not a decryption failure; warning on it would train operators to ignore the log. + import logging + + prisma = _make_prisma_with_existing(row=_legacy_row("sk-live-byok-secret")) + + with caplog.at_level(logging.WARNING, logger="LiteLLM Proxy"): + assert await get_user_oauth_credential(prisma, "alice", "srv-1") is None + + assert [rec.getMessage() for rec in caplog.records if "could not be decrypted" in rec.getMessage()] == [] + + # ── list_user_oauth_credentials ─────────────────────────────────────────────── diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index bdaf1458fb0..34852850de6 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -4,6 +4,7 @@ import hashlib import json import time from base64 import urlsafe_b64encode +from typing import TYPE_CHECKING from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -11,6 +12,11 @@ from fastapi import HTTPException from litellm.types.mcp import MCPAuth +if TYPE_CHECKING: + import httpx + + from litellm.types.mcp_server.mcp_server_manager import MCPServer + # Fixture to mock IP address check for all MCP tests # This prevents tests from failing due to IP-based access control diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py index 7bdd3b36763..feab179570b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_stale_session.py @@ -1411,6 +1411,8 @@ async def _run_passthrough_connect( ): """Drive handle_streamable_http_mcp through the preemptive-401 gate and report whether it challenged (raised) or forwarded to the session manager. Returns (challenged, www_authenticate).""" + from fastapi import HTTPException + from litellm.proxy._experimental.mcp_server.server import ( handle_streamable_http_mcp, session_manager_stateless, diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 7ba9f463197..d7fb121ef9b 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -1,9 +1,13 @@ import asyncio import json +import sys from datetime import datetime from typing import Any, Dict, Optional from unittest.mock import AsyncMock, MagicMock +if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11 + from exceptiongroup import BaseExceptionGroup + import httpx import pytest from fastapi import HTTPException diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py index 83e4dcf5677..da7f43c7118 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_semantic_tool_filter.py @@ -12,6 +12,9 @@ from unittest.mock import AsyncMock, Mock, patch import pytest +if sys.version_info < (3, 11): # BaseExceptionGroup is a builtin only from 3.11 + from exceptiongroup import BaseExceptionGroup + sys.path.insert(0, os.path.abspath("../..")) from mcp.types import Tool as MCPTool diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py index e54358c1f00..2ff38af80b1 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_endpoints.py @@ -1317,15 +1317,373 @@ async def test_handle_stream_message_rejects_invalid_params_with_32602(): request_id="req-1", params={"message": 12345}, ) + assert response.media_type == "text/event-stream" chunks = [chunk async for chunk in response.body_iterator] body = "".join( chunk.decode() if isinstance(chunk, bytes) else chunk for chunk in chunks ) - payload = json.loads(body.strip()) + assert body.startswith("data: ") + assert body.endswith("\n\n") + payload = json.loads(body.removeprefix("data: ").strip()) assert payload["error"]["code"] == -32602 assert payload["id"] == "req-1" +@pytest.mark.asyncio +async def test_handle_stream_message_frames_events_as_sse(): + """message/stream must return text/event-stream with each JSON-RPC object + framed as ``data: \\n\\n``. Regression for #35027: NDJSON framing + breaks the official a2a-sdk client, which requires SSE.""" + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + + events = [ + { + "jsonrpc": "2.0", + "id": "req-1", + "result": {"kind": "task", "id": "t-1", "status": {"state": "working"}}, + }, + { + "jsonrpc": "2.0", + "id": "req-1", + "result": {"kind": "message", "parts": [{"kind": "text", "text": "pong"}]}, + }, + ] + + async def fake_stream(**kwargs): + for event in events: + yield event + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch( + "litellm.a2a_protocol.asend_message_streaming", + new=fake_stream, + ) + ) + + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + } + }, + ) + + assert response.media_type == "text/event-stream" + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert len(chunks) == len(events) + for chunk, event in zip(chunks, events): + assert chunk.startswith("data: ") + assert chunk.endswith("\n\n") + assert json.loads(chunk.removeprefix("data: ").strip()) == event + + +@pytest.mark.asyncio +async def test_handle_stream_message_sdk_unavailable_frames_error_as_sse(): + """When the a2a package is unavailable the -32603 error must still be + emitted as a single SSE event so the a2a-sdk client can parse it.""" + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + + with patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", False): + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={"message": {"role": "user", "parts": []}}, + ) + + assert response.media_type == "text/event-stream" + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + assert len(chunks) == 1 + assert chunks[0].startswith("data: ") + assert chunks[0].endswith("\n\n") + payload = json.loads(chunks[0].removeprefix("data: ").strip()) + assert payload["error"]["code"] == -32603 + assert payload["id"] == "req-1" + + +@pytest.mark.asyncio +async def test_handle_stream_message_proxy_hook_path_frames_events_as_sse(): + """When proxy hooks are wired the events are routed through + async_streaming_data_generator; that path must also frame each JSON-RPC + object as ``data: \\n\\n`` (regression for #35027).""" + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + from litellm.proxy.utils import ProxyLogging + + events = [ + {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}}, + { + "jsonrpc": "2.0", + "id": "req-1", + "result": {"kind": "message", "parts": [{"kind": "text", "text": "pong"}]}, + }, + ] + + async def fake_stream(**kwargs): + for event in events: + yield event + + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream) + ) + + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + } + }, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + request_data={"model": "a2a/test"}, + proxy_logging_obj=proxy_logging_obj, + ) + + assert response.media_type == "text/event-stream" + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert len(chunks) == len(events) + for chunk, event in zip(chunks, events): + assert chunk.startswith("data: ") + assert chunk.endswith("\n\n") + assert json.loads(chunk.removeprefix("data: ").strip()) == event + + +@pytest.mark.asyncio +async def test_handle_stream_message_frames_preserialized_jsonrpc_error_once(): + """A stream chunk that is already a serialized JSON-RPC object (what a + guardrail may yield when it terminates an A2A stream mid-flight) must be + framed as one SSE event carrying that object, not JSON-encoded a second time + into a bare string.""" + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + + error_event = { + "jsonrpc": "2.0", + "id": "req-1", + "error": {"code": -32603, "message": "blocked by guardrail", "data": {}}, + } + + async def fake_stream(**kwargs): + yield json.dumps(error_event) + "\n" + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream) + ) + + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + } + }, + ) + + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert len(chunks) == 1 + payload = json.loads(chunks[0].removeprefix("data: ").strip()) + assert payload == error_event + + +@pytest.mark.asyncio +async def test_handle_stream_message_proxy_hook_path_frames_errors_as_sse(): + """A failure while the hooked generator is streaming must reach the client as + a ``data:``-framed JSON-RPC error, not as a bare NDJSON line.""" + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + from litellm.proxy.utils import ProxyLogging + + async def fake_stream(**kwargs): + yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}} + raise ValueError("upstream died") + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream) + ) + + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + } + }, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + request_data={"model": "a2a/test"}, + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + ) + + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert len(chunks) == 2 + assert chunks[-1].startswith("data: ") + error_payload = json.loads(chunks[-1].removeprefix("data: ").strip()) + assert error_payload["id"] == "req-1" + assert error_payload["error"]["code"] == -32603 + assert "upstream died" in error_payload["error"]["message"] + + +@pytest.mark.asyncio +async def test_handle_stream_message_frames_upstream_call_failure_as_sse_error(): + """A failure raised before any event is streamed (with proxy hooks wired) is + still delivered as a ``data:``-framed JSON-RPC error.""" + from litellm.caching.caching import DualCache + from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + from litellm.proxy.utils import ProxyLogging + + def fake_stream(**kwargs): + raise ValueError("could not reach agent") + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream) + ) + + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + } + }, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + request_data={"model": "a2a/test"}, + proxy_logging_obj=ProxyLogging(user_api_key_cache=DualCache()), + ) + + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert len(chunks) == 1 + error_payload = json.loads(chunks[0].removeprefix("data: ").strip()) + assert error_payload["id"] == "req-1" + assert error_payload["error"]["code"] == -32603 + assert "could not reach agent" in error_payload["error"]["message"] + + +@pytest.mark.asyncio +async def test_handle_stream_message_forwards_unparseable_chunk_as_sse_event(): + """A chunk that is not JSON at all still leaves as one well-formed SSE event + instead of raising and killing the stream.""" + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + + async def fake_stream(**kwargs): + yield "not json at all" + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream) + ) + + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + } + }, + ) + + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert chunks == ['data: "not json at all"\n\n'] + + +@pytest.mark.asyncio +async def test_handle_stream_message_frames_mid_stream_failure_as_sse_error(): + """An upstream failure after the response started is reported as a + ``data:``-framed JSON-RPC error object, so an SSE client sees the failure.""" + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + + async def fake_stream(**kwargs): + yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}} + raise RuntimeError("upstream died") + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream) + ) + + response = await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + } + }, + ) + + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert len(chunks) == 2 + error_payload = json.loads(chunks[-1].removeprefix("data: ").strip()) + assert error_payload["id"] == "req-1" + assert error_payload["error"]["code"] == -32603 + assert "upstream died" in error_payload["error"]["message"] + + @pytest.mark.asyncio async def test_send_message_pascal_case_routes_to_asend_message(): from litellm.proxy._types import UserAPIKeyAuth @@ -2029,3 +2387,83 @@ async def test_forward_jsonrpc_sse_is_untouched_while_keepalives_are_unconfigure assert not any(chunk.startswith(":") for chunk in chunks) assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task" + + +async def _stream_message_response(): + from litellm.proxy.agent_endpoints.a2a_endpoints import _handle_stream_message + + return await _handle_stream_message( + api_base="http://upstream.local", + request_id="req-1", + params={ + "message": { + "role": "user", + "parts": [{"kind": "text", "text": "hi"}], + "messageId": "msg-1", + } + }, + ) + + +@pytest.mark.asyncio +async def test_handle_stream_message_pings_while_the_upstream_agent_is_still_silent( + monkeypatch, +): + """message/stream is SSE like tasks/resubscribe, so a slow first event must be + held open by the same keepalives rather than sitting idle for the whole + time-to-first-token.""" + import asyncio + + import litellm + + monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", 0.05) + + async def fake_stream(**kwargs): + await asyncio.sleep(0.3) + yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}} + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream) + ) + + response = await _stream_message_response() + assert response.headers["x-accel-buffering"] == "no" + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert chunks[0] == ": ping\n\n" + assert chunks.count(": ping\n\n") >= 3 + assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task" + + +@pytest.mark.asyncio +async def test_handle_stream_message_is_untouched_while_keepalives_are_unconfigured( + monkeypatch, +): + """Off until an operator sets an interval, so the default stream is unchanged.""" + import litellm + + monkeypatch.setattr(litellm, "sse_keepalive_ping_interval_seconds", None) + + async def fake_stream(**kwargs): + yield {"jsonrpc": "2.0", "id": "req-1", "result": {"kind": "task", "id": "t-1"}} + + with ExitStack() as stack: + stack.enter_context(patch("litellm.a2a_protocol.main.A2A_SDK_AVAILABLE", True)) + stack.enter_context( + patch("litellm.a2a_protocol.asend_message_streaming", new=fake_stream) + ) + + response = await _stream_message_response() + assert "x-accel-buffering" not in response.headers + chunks = [ + chunk.decode() if isinstance(chunk, bytes) else chunk + async for chunk in response.body_iterator + ] + + assert not any(chunk.startswith(":") for chunk in chunks) + assert json.loads(chunks[-1].removeprefix("data: "))["result"]["kind"] == "task" diff --git a/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py b/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py index 069c72af53a..3dc4d3427cd 100644 --- a/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py +++ b/tests/test_litellm/proxy/agent_endpoints/test_a2a_version_e2e.py @@ -304,6 +304,7 @@ async def test_proxy_streaming_serves_1_0_envelopes(): user_api_key_dict=user_api_key_dict, ) + assert response.media_type == "text/event-stream" lines: List[Dict[str, Any]] = [] async for raw_line in response.body_iterator: line = ( @@ -312,7 +313,7 @@ async def test_proxy_streaming_serves_1_0_envelopes(): else str(raw_line).strip() ) if line: - lines.append(json.loads(line)) + lines.append(json.loads(line.removeprefix("data:").strip())) assert lines, "expected at least one streamed JSON-RPC event" message_events = [ diff --git a/tests/test_litellm/proxy/auth/test_auth_checks.py b/tests/test_litellm/proxy/auth/test_auth_checks.py index 1d1bd9ebf8a..ccbf00a67b9 100644 --- a/tests/test_litellm/proxy/auth/test_auth_checks.py +++ b/tests/test_litellm/proxy/auth/test_auth_checks.py @@ -13,7 +13,7 @@ from datetime import datetime, timedelta, timezone import httpx import pytest -from fastapi import status +from fastapi import Request, status import litellm from litellm.proxy._types import ( @@ -2760,11 +2760,9 @@ async def test_common_checks_metadata_route_keeps_key_tags_out_of_provider_metad assert "metadata" not in request_body -def _pass_through_request() -> "Request": +def _pass_through_request() -> Request: """A Request whose FastAPI-resolved endpoint carries the pass-through marker, i.e. the request was dispatched to a user-defined pass-through handler.""" - from fastapi import Request - from litellm.types.passthrough_endpoints.pass_through_endpoints import ( LITELLM_PASS_THROUGH_ENDPOINT_MARKER, ) @@ -2776,10 +2774,9 @@ def _pass_through_request() -> "Request": return Request(scope={"type": "http", "headers": [], "endpoint": pass_through_endpoint}) -def _builtin_request() -> "Request": +def _builtin_request() -> Request: """A Request dispatched to a built-in (non-pass-through) handler, e.g. what a custom path colliding with a core route actually resolves to.""" - from fastapi import Request def chat_completions(): ... @@ -5698,6 +5695,64 @@ async def test_get_default_end_user_budget_db_fetch_returns_validated_budget(mon assert mock_cache.async_set_cache.call_args.kwargs["value"] is result +@pytest.mark.asyncio +async def test_get_team_member_default_budget_caches_json_safe_payload(): + """The Redis layer json.dumps() the cached value, so datetime columns on the budget row + must be dumped to ISO strings before the write, and the read side must give back a model. + """ + from litellm.proxy.auth.auth_checks import get_team_member_default_budget + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + budget_row = MagicMock() + budget_row.dict = lambda: { + "budget_id": "tm-budget-1", + "max_budget": 25.0, + "created_at": datetime(2026, 1, 1, tzinfo=timezone.utc), + "updated_at": datetime(2026, 1, 2, tzinfo=timezone.utc), + } + + mock_prisma_client = MagicMock() + mock_prisma_client.db.litellm_budgettable.find_unique = AsyncMock(return_value=budget_row) + + class _JsonOnlyRedis: + """Stands in for RedisCache, which serializes with a bare json.dumps().""" + + def __init__(self): + self.writes = [] + + async def async_set_cache(self, key, value, **kwargs): + self.writes.append((key, json.dumps(value))) + + async def async_get_cache(self, key, **kwargs): + return None + + redis_cache = _JsonOnlyRedis() + cache = UserApiKeyCache(redis_cache=redis_cache) + + budget = await get_team_member_default_budget( + budget_id="tm-budget-1", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + ) + + assert isinstance(budget, LiteLLM_BudgetTable) + assert budget.max_budget == 25.0 + assert len(redis_cache.writes) == 1 + written_key, written_payload = redis_cache.writes[0] + assert written_key == "team_member_default_budget:tm-budget-1" + assert json.loads(written_payload)["max_budget"] == 25.0 + + cached = await get_team_member_default_budget( + budget_id="tm-budget-1", + prisma_client=mock_prisma_client, + user_api_key_cache=cache, + ) + + assert isinstance(cached, LiteLLM_BudgetTable) + assert cached.max_budget == 25.0 + mock_prisma_client.db.litellm_budgettable.find_unique.assert_awaited_once() + + @pytest.mark.asyncio async def test_get_end_user_object_db_fetch_returns_validated_end_user(): from litellm.proxy.auth.auth_checks import get_end_user_object diff --git a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py index 27798ec0bff..b4725a81823 100644 --- a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py +++ b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py @@ -31,6 +31,7 @@ sys.path.insert( ) # Adds the parent directory to the system path from litellm._logging import verbose_proxy_logger +from litellm.exceptions import BudgetExceededError from litellm.proxy._types import ProxyErrorTypes, ProxyException, UserAPIKeyAuth from litellm.proxy.auth.auth_exception_handler import UserAPIKeyAuthExceptionHandler @@ -511,3 +512,196 @@ async def test_auth_failure_without_resolved_identity_still_logs(): assert logged.api_key != "sk-unknown" assert logged.api_key == UserAPIKeyAuth(api_key="sk-unknown").api_key assert logged.request_route == "/v1/chat/completions" + + +def _http_request(client_host: str | None = "10.1.2.3", headers: dict[str, str] | None = None) -> Request: + return Request( + { + "type": "http", + "http_version": "1.1", + "method": "POST", + "scheme": "http", + "path": "/v1/chat/completions", + "raw_path": b"/v1/chat/completions", + "query_string": b"", + "root_path": "", + "server": ("testserver", 80), + "client": (client_host, 51234) if client_host is not None else None, + "headers": [(k.lower().encode(), v.encode()) for k, v in (headers or {}).items()], + } + ) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "auth_error, general_settings, request_kwargs, expected_ip", + [ + pytest.param( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + {"allow_requests_on_db_unavailable": False}, + {}, + "10.1.2.3", + id="401_socket_peer", + ), + pytest.param( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + {"allow_requests_on_db_unavailable": False, "use_x_forwarded_for": True}, + {"headers": {"x-forwarded-for": "203.0.113.9"}}, + "203.0.113.9", + id="401_x_forwarded_for", + ), + pytest.param( + BudgetExceededError(message="Budget exceeded", current_cost=100, max_budget=100), + {"allow_requests_on_db_unavailable": False}, + {}, + "10.1.2.3", + id="429_budget_exceeded", + ), + ], +) +async def test_auth_failure_logs_requester_ip_address( + auth_error: Exception, + general_settings: dict[str, bool], + request_kwargs: dict[str, dict[str, str]], + expected_ip: str, +) -> None: + """401s and budget 429s are rejected before `add_litellm_data_to_request` stamps + the caller IP, so without this the failure logs (spend logs, prometheus client_ip) + had no IP, and a 401 rarely carries a key or user identity either.""" + with ( + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ) as mock_hook, + patch("litellm.proxy.proxy_server.general_settings", general_settings), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + auth_error, + _http_request(**request_kwargs), + {"model": "gpt-4o"}, + "/v1/chat/completions", + None, + "sk-bad-key", + ) + + logged_request_data = mock_hook.call_args[1]["request_data"] + assert logged_request_data["metadata"]["requester_ip_address"] == expected_ip + + +@pytest.mark.asyncio +async def test_auth_failure_keeps_existing_requester_ip_address(): + """An IP already recorded upstream (e.g. a trusted-proxy resolved value) wins over + the socket peer.""" + with ( + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ) as mock_hook, + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + _http_request(), + {"metadata": {"requester_ip_address": "198.51.100.4"}}, + "/v1/chat/completions", + None, + "sk-bad-key", + ) + + logged_request_data = mock_hook.call_args[1]["request_data"] + assert logged_request_data["metadata"]["requester_ip_address"] == "198.51.100.4" + + +@pytest.mark.asyncio +async def test_auth_failure_ip_uses_litellm_metadata_when_present(): + """Routes that keep proxy metadata under `litellm_metadata` (e.g. /responses) must + get the IP there, since that is the dict the logging layer reads for them.""" + with ( + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ) as mock_hook, + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + _http_request(), + {"litellm_metadata": {}, "metadata": {"user_supplied": "keep-me"}}, + "/v1/responses", + None, + "sk-bad-key", + ) + + logged_request_data = mock_hook.call_args[1]["request_data"] + assert logged_request_data["litellm_metadata"]["requester_ip_address"] == "10.1.2.3" + assert logged_request_data["metadata"] == {"user_supplied": "keep-me"} + + +@pytest.mark.asyncio +async def test_auth_failure_ip_stamp_does_not_mutate_callers_request_data(): + """The handler must not rewrite the caller's dict; the IP is for the failure log only.""" + request_data = {"model": "gpt-4o"} + + with ( + patch("litellm.proxy.auth.auth_exception_handler.seed_request_identity"), + patch( + "litellm.proxy.proxy_server.proxy_logging_obj.post_call_failure_hook", + new_callable=AsyncMock, + return_value=None, + ), + patch( + "litellm.proxy.proxy_server.general_settings", + {"allow_requests_on_db_unavailable": False}, + ), + ): + with pytest.raises(ProxyException): + await UserAPIKeyAuthExceptionHandler._handle_authentication_error( + ProxyException( + message="Invalid API key", + type=ProxyErrorTypes.auth_error, + param=None, + code=status.HTTP_401_UNAUTHORIZED, + ), + _http_request(), + request_data, + "/v1/chat/completions", + None, + "sk-bad-key", + ) + + assert request_data == {"model": "gpt-4o"} diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index 3102e69bf26..ecf7f89d487 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -2226,6 +2226,66 @@ class TestIsRequestBodySafeBlocksBedrockProjectOverride: ) +class TestIsRequestBodySafeBlocksRustOptIn: + """``rust`` hands the whole call to the Rust core, which signs and sends + with its own HTTP client rather than the one the deployment configured, and + reports no ``post_call``. The proxy splats the request body straight into + the router, and ``rust`` is a litellm param, so it lands in + ``litellm_params`` and the gate honours it: without this entry any + authenticated caller picks a transport and a callback surface the admin + never chose. It stays a deployment decision, liftable only by the same + admin opt-in as the rest of the list.""" + + def test_rust_in_request_body_is_rejected(self): + with pytest.raises(ValueError, match="rust"): + is_request_body_safe( + request_body={"model": "gpt-4", "rust": True}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_rust_under_extra_body_is_rejected(self): + with pytest.raises(ValueError, match="not allowed in request body"): + is_request_body_safe( + request_body={"model": "gpt-4", "extra_body": {"rust": True}}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_api_key_does_not_bypass_the_rust_block(self): + with pytest.raises(ValueError, match="rust"): + is_request_body_safe( + request_body={"model": "gpt-4", "api_key": "sk-anything", "rust": True}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + + def test_admin_opt_in_proxy_wide_allows_rust(self): + assert ( + is_request_body_safe( + request_body={"model": "gpt-4", "rust": True}, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="gpt-4", + ) + is True + ) + + def test_body_without_rust_is_still_allowed(self): + assert ( + is_request_body_safe( + request_body={"model": "gpt-4", "temperature": 0.7}, + general_settings={}, + llm_router=None, + model="gpt-4", + ) + is True + ) + + class TestIsRequestBodySafeBlocksVertexCredentialAlias: @pytest.mark.parametrize("field", ["vertex_ai_credentials"]) def test_field_in_request_body_is_rejected(self, field): diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 3840c90d691..a9e12beb54b 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -1,10 +1,16 @@ +import asyncio +import re +import time +from collections.abc import Mapping, Sequence from typing import Optional from unittest.mock import AsyncMock, MagicMock, patch from fastapi import HTTPException +import httpx import pytest from litellm.proxy._types import ( + DEFAULT_JWKS_STALE_TTL, JWTLiteLLMRoleMap, LiteLLM_JWTAuth, LiteLLM_TeamMembership, @@ -15,7 +21,16 @@ from litellm.proxy._types import ( ProxyErrorTypes, ProxyException, ) -from litellm.proxy.auth.handle_jwt import JWTAuthManager, JWTHandler +from litellm.caching.dual_cache import DualCache +from litellm.proxy.auth.handle_jwt import ( + JWKS_FETCH_ATTEMPTS, + STALE_CACHE_KEY_PREFIX, + STALE_WRITTEN_AT_CACHE_KEY_PREFIX, + JWKSUnreachableError, + JWTAuthManager, + JWTHandler, + NoMatchingJWTPublicKeyError, +) @pytest.mark.asyncio @@ -3921,6 +3936,7 @@ async def test_get_public_key_fetches_and_caches_jwks_response(): expected_key_id = "cached-key" _, jwk = _get_rsa_key_and_jwk(kid=expected_key_id) mock_response = MagicMock() + mock_response.status_code = 200 mock_response.json.return_value = {"keys": [jwk]} jwt_handler.http_handler.get = AsyncMock(return_value=mock_response) @@ -3936,6 +3952,560 @@ async def test_get_public_key_fetches_and_caches_jwks_response(): assert cached_keys == [jwk] +class _ScriptedJWKSEndpoint: + """Injected stand-in for ``JWTHandler.http_handler`` with scripted per-call outcomes. + + Each outcome is either an exception to raise or a JSON body to return; the + last outcome repeats for any further calls. + """ + + def __init__( + self, + outcomes: Sequence[Exception | Mapping[str, object] | MagicMock], + delay: float = 0.0, + ) -> None: + self.outcomes = outcomes + self.delay = delay + self.call_count = 0 + + async def get( + self, + url: str, + params: Mapping[str, str] | None = None, + headers: Mapping[str, str] | None = None, + ) -> MagicMock: + self.call_count += 1 + if self.delay: + await asyncio.sleep(self.delay) + outcome = self.outcomes[min(self.call_count - 1, len(self.outcomes) - 1)] + if isinstance(outcome, Exception): + raise outcome + if isinstance(outcome, MagicMock): + return outcome + response = MagicMock() + response.status_code = 200 + response.json.return_value = outcome + return response + + +def _get_jwt_handler_with_scripted_endpoint( + cache: "DualCache", + endpoint: _ScriptedJWKSEndpoint, + public_key_ttl: float = 600, + public_key_stale_ttl: float = DEFAULT_JWKS_STALE_TTL, +) -> JWTHandler: + jwt_handler = JWTHandler() + jwt_handler.update_environment( + prisma_client=None, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth( + public_key_ttl=public_key_ttl, + public_key_stale_ttl=public_key_stale_ttl, + ), + ) + jwt_handler.http_handler = endpoint + return jwt_handler + + +@pytest.mark.asyncio +async def test_get_public_key_retries_transient_jwks_fetch_failure(): + """A single connect timeout to the IdP must be retried, not surfaced to the caller.""" + from litellm.caching.dual_cache import DualCache + + _, jwk = _get_rsa_key_and_jwk(kid="retried-key") + endpoint = _ScriptedJWKSEndpoint((httpx.ConnectTimeout("connect timed out"), {"keys": [jwk]})) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(DualCache(), endpoint) + + public_key = await jwt_handler._get_public_key_from_jwks_url( + jwks_url="https://issuer.example.com/keys", + kid="retried-key", + ) + + assert public_key == jwk + assert endpoint.call_count == 2 + + +@pytest.mark.asyncio +async def test_get_public_key_serves_stale_keys_when_jwks_refresh_fails(): + """Once the TTL lapses, an unreachable IdP must not invalidate a still-valid signing key.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://issuer.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="stale-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="stale-key") == jwk + + await cache.async_delete_cache(key=f"litellm_jwt_auth_keys_{jwks_url}") + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + + public_key = await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="stale-key") + + assert public_key == jwk + + +@pytest.mark.asyncio +async def test_stale_jwks_window_is_the_configured_grace_past_a_long_public_key_ttl(): + """The stale window is `public_key_stale_ttl` past the active entry, whatever `public_key_ttl` is set to. + + Deriving the window from `public_key_ttl` instead would collapse it to nothing on the long TTLs that + make the fallback worth having. + """ + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://long-ttl-issuer.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="long-ttl-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint( + cache, + endpoint, + public_key_ttl=90000, + public_key_stale_ttl=3600, + ) + + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="long-ttl-key") + + active_key = f"litellm_jwt_auth_keys_{jwks_url}" + active_deadline = cache.in_memory_cache.ttl_dict[active_key] + stale_deadline = cache.in_memory_cache.ttl_dict[f"{STALE_CACHE_KEY_PREFIX}{active_key}"] + + assert stale_deadline - active_deadline == pytest.approx(3600, abs=1) + + +@pytest.mark.asyncio +async def test_long_public_key_ttl_still_serves_stale_keys_when_the_idp_is_unreachable(): + """A long `public_key_ttl` must not leave the stale fallback inert once that TTL finally lapses.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://long-ttl-fallback.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="long-ttl-fallback-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_ttl=604800) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="long-ttl-fallback-key") == jwk + + await cache.async_delete_cache(key=f"litellm_jwt_auth_keys_{jwks_url}") + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="long-ttl-fallback-key") == jwk + + +@pytest.mark.asyncio +async def test_removed_signing_key_stops_being_trusted_once_the_stale_window_expires(monkeypatch): + """The stale fallback is bounded: past its window a key the IdP dropped is no longer served.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://revoking-issuer.example.com/keys" + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", jwks_url) + + _, jwk = _get_rsa_key_and_jwk(kid="revoked-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + assert await jwt_handler.get_public_key(kid="revoked-key") == jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + + assert await jwt_handler.get_public_key(kid="revoked-key") == jwk + + await cache.async_delete_cache(key=f"{STALE_CACHE_KEY_PREFIX}{active_cache_key}") + + with pytest.raises(ProxyException) as exc_info: + await jwt_handler.get_public_key(kid="revoked-key") + + assert exc_info.value.code == "503" + assert exc_info.value.type == ProxyErrorTypes.auth_provider_unavailable + + +@pytest.mark.asyncio +async def test_key_removed_from_a_reachable_jwks_is_rejected_without_consulting_the_stale_copy(): + """A reachable IdP always wins: dropping a key revokes it immediately, stale copy included.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://rotating-issuer.example.com/keys" + _, retired_jwk = _get_rsa_key_and_jwk(kid="retired-key") + _, current_jwk = _get_rsa_key_and_jwk(kid="current-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [retired_jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="retired-key") == retired_jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + endpoint.outcomes = ({"keys": [current_jwk]},) + + with pytest.raises(NoMatchingJWTPublicKeyError): + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="retired-key") + + assert await cache.async_get_cache(key=f"{STALE_CACHE_KEY_PREFIX}{active_cache_key}") == [current_jwk] + + +@pytest.mark.asyncio +async def test_zero_public_key_stale_ttl_fails_closed_instead_of_serving_stale_keys(): + """`public_key_stale_ttl=0` is the escape hatch for deployments that cannot trust an unrefreshed key.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://fail-closed-issuer.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="fail-closed-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_stale_ttl=0) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="fail-closed-key") == jwk + assert await cache.async_get_cache(key=f"{STALE_CACHE_KEY_PREFIX}litellm_jwt_auth_keys_{jwks_url}") is None + + await cache.async_delete_cache(key=f"litellm_jwt_auth_keys_{jwks_url}") + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + + with pytest.raises(JWKSUnreachableError): + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="fail-closed-key") + + +@pytest.mark.asyncio +@pytest.mark.parametrize("lowered_stale_ttl", [0, 30]) +async def test_lowering_public_key_stale_ttl_stops_serving_a_copy_cached_under_the_old_setting(lowered_stale_ttl): + """Lowering the window has to bite immediately: an operator does this mid-incident, on a shared cache. + + The stale entry keeps whatever expiry it was written with, so enforcing the bound only at write time would + leave a copy taken under the old, longer setting servable until it aged out on its own. + """ + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://relaxed-then-tightened.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="tightened-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + generous = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_stale_ttl=86400) + + assert await generous._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="tightened-key") == jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + assert await cache.async_get_cache(key=f"{STALE_CACHE_KEY_PREFIX}{active_cache_key}") == [jwk] + + # The operator tightens the window and restarts; the cache, and its long-lived copy, survive. + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + tightened = _get_jwt_handler_with_scripted_endpoint( + cache, endpoint, public_key_stale_ttl=lowered_stale_ttl + ) + await cache.async_set_cache( + key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}", + value=time.time() - 7200, + ttl=86400, + ) + + with pytest.raises(JWKSUnreachableError): + await tightened._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="tightened-key") + + +@pytest.mark.asyncio +async def test_zero_public_key_stale_ttl_fails_closed_even_for_a_freshly_written_copy(): + """`0` must fail closed on its own, not merely because the copy happens to be older than `public_key_ttl`. + + The active entry can disappear before it expires, through cache eviction or a flush, which leaves a stale + copy younger than `public_key_ttl`. Bounding only on age would still serve it. + """ + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://evicted-active-entry.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="fresh-copy-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + generous = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_ttl=600, public_key_stale_ttl=3600) + + assert await generous._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="fresh-copy-key") == jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + written_at = await cache.async_get_cache(key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}") + assert time.time() - written_at < 600 + + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + fail_closed = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_ttl=600, public_key_stale_ttl=0) + + with pytest.raises(JWKSUnreachableError): + await fail_closed._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="fresh-copy-key") + + +@pytest.mark.asyncio +async def test_stale_copy_with_no_recorded_write_time_is_not_served(): + """The bound is enforced from the recorded write time, so losing it must fail closed, never open.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://undated-copy.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="undated-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="undated-key") == jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + await cache.async_delete_cache(key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}") + assert await cache.async_get_cache(key=f"{STALE_CACHE_KEY_PREFIX}{active_cache_key}") == [jwk] + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + + with pytest.raises(JWKSUnreachableError): + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="undated-key") + + +@pytest.mark.asyncio +async def test_increasing_public_key_stale_ttl_only_extends_within_the_new_bound(): + """Raising the window re-measures from the copy's refresh time; it does not bless whatever is cached.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://widened-window.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="widened-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + narrow = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_ttl=600, public_key_stale_ttl=60) + + assert await narrow._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="widened-key") == jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + written_at_key = f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}" + await cache.async_delete_cache(key=active_cache_key) + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + widened = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_ttl=600, public_key_stale_ttl=3600) + + # Older than the widened bound of 600 + 3600, so widening must not revive it. + await cache.async_set_cache(key=written_at_key, value=time.time() - 5000, ttl=86400) + with pytest.raises(JWKSUnreachableError): + await widened._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="widened-key") + + # Inside the widened bound, so it is servable again. + await cache.async_set_cache(key=written_at_key, value=time.time() - 1000, ttl=86400) + assert await widened._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="widened-key") == jwk + + +@pytest.mark.asyncio +async def test_stale_copy_written_at_survives_a_whole_number_epoch(): + """A Redis JSON round-trip can return the epoch as an int, and that must not read as a missing timestamp. + + Rejecting it would fail closed on a copy that is well inside the window, in the shared-cache deployment + the stale fallback exists to serve. + """ + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://int-epoch.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="int-epoch-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="int-epoch-key") == jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + await cache.async_set_cache( + key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}", + value=int(time.time()) - 60, + ttl=86400, + ) + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="int-epoch-key") == jwk + + +@pytest.mark.asyncio +async def test_stale_copy_with_a_malformed_write_time_is_not_served(): + """An unreadable refresh timestamp is indistinguishable from an unbounded one, so it fails closed.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://malformed-timestamp.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="malformed-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="malformed-key") == jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + await cache.async_set_cache( + key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}", + value="whenever", + ttl=86400, + ) + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + + with pytest.raises(JWKSUnreachableError): + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="malformed-key") + + +@pytest.mark.asyncio +async def test_public_key_stale_ttl_defaults_to_one_hour(): + """The default is the exposure bound for a key the IdP revoked mid-outage, so it stays short deliberately.""" + assert LiteLLM_JWTAuth().public_key_stale_ttl == 3600 + + +@pytest.mark.asyncio +async def test_stale_fallback_warns_with_the_kid_and_how_stale_the_jwks_copy_is(caplog): + """Serving an unrefreshed signing key is a security-relevant event, so it must be legible in the logs.""" + import logging + + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://warned-issuer.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="warned-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint, public_key_stale_ttl=1800) + + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="warned-key") + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + await cache.async_set_cache( + key=f"{STALE_WRITTEN_AT_CACHE_KEY_PREFIX}{active_cache_key}", + value=time.time() - 120, + ttl=600, + ) + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + + caplog.set_level(logging.WARNING) + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="warned-key") + + warnings = [r.getMessage() for r in caplog.records if r.levelno == logging.WARNING] + stale_warnings = [m for m in warnings if "stale JWKS copy" in m] + assert len(stale_warnings) == 1 + assert "kid=warned-key" in stale_warnings[0] + assert jwks_url in stale_warnings[0] + + freshness = re.search(r"last refreshed (\d+)s ago, stops being trusted in (\d+)s", stale_warnings[0]) + assert freshness is not None + age, remaining = int(freshness.group(1)), int(freshness.group(2)) + assert age == pytest.approx(120, abs=2) + assert remaining == pytest.approx(600 + 1800 - 120, abs=2) + + +@pytest.mark.asyncio +async def test_unparseable_jwks_response_does_not_fall_back_to_the_stale_copy(): + """Only an unreachable IdP unlocks the stale copy. A reachable one that answers badly must surface the error.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://garbled-issuer.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="garbled-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="garbled-key") == jwk + + await cache.async_delete_cache(key=f"litellm_jwt_auth_keys_{jwks_url}") + garbled = MagicMock() + garbled.status_code = 200 + garbled.text = "not json" + garbled.json.side_effect = ValueError("Expecting value: line 1 column 1") + endpoint.outcomes = (garbled,) + + with pytest.raises(Exception, match="Error parsing response"): + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="garbled-key") + + +@pytest.mark.asyncio +async def test_jwks_error_response_is_not_cached_over_the_last_known_good_keys(): + """An IdP error body must never be stored as the key set, least of all as the stale copy.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://erroring-issuer.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="erroring-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + assert await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="erroring-key") == jwk + + active_cache_key = f"litellm_jwt_auth_keys_{jwks_url}" + await cache.async_delete_cache(key=active_cache_key) + server_error = MagicMock() + server_error.status_code = 503 + server_error.text = '{"error": "upstream unavailable"}' + server_error.json.return_value = {"error": "upstream unavailable"} + endpoint.outcomes = (server_error,) + + with pytest.raises(Exception, match="returned status 503"): + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="erroring-key") + + assert await cache.async_get_cache(key=active_cache_key) is None + assert await cache.async_get_cache(key=f"{STALE_CACHE_KEY_PREFIX}{active_cache_key}") == [jwk] + + +@pytest.mark.asyncio +async def test_sustained_jwks_outage_refetches_once_per_backoff_window_not_once_per_request(): + """Without a backoff, every request during an outage pays three timeouts serialised behind the refresh lock.""" + from litellm.caching.dual_cache import DualCache + + jwks_url = "https://flooded-issuer.example.com/keys" + _, jwk = _get_rsa_key_and_jwk(kid="flooded-key") + cache = DualCache() + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(cache, endpoint) + + await jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="flooded-key") + + await cache.async_delete_cache(key=f"litellm_jwt_auth_keys_{jwks_url}") + endpoint.outcomes = (httpx.ConnectTimeout("connect timed out"),) + calls_before_outage = endpoint.call_count + + public_keys = await asyncio.gather( + *[jwt_handler._get_public_key_from_jwks_url(jwks_url=jwks_url, kid="flooded-key") for _ in range(6)] + ) + + assert public_keys == [jwk] * 6 + assert endpoint.call_count - calls_before_outage == JWKS_FETCH_ATTEMPTS + + +@pytest.mark.asyncio +async def test_get_public_key_raises_503_when_jwks_unreachable_and_no_cached_keys(monkeypatch): + """An unreachable IdP is an infra failure: 503, never a 401 that clients read as bad credentials.""" + from litellm.caching.dual_cache import DualCache + + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", "https://issuer.example.com/keys") + endpoint = _ScriptedJWKSEndpoint((httpx.ConnectTimeout("connect timed out"),)) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(DualCache(), endpoint) + + with pytest.raises(ProxyException) as exc_info: + await jwt_handler.get_public_key(kid="any-key") + + assert exc_info.value.code == "503" + assert exc_info.value.type == ProxyErrorTypes.auth_provider_unavailable + assert "ConnectTimeout" in exc_info.value.message + assert endpoint.call_count == 3 + + +@pytest.mark.asyncio +async def test_get_public_key_coalesces_concurrent_jwks_refreshes(): + """Concurrent requests in the TTL-expiry window share one JWKS fetch.""" + from litellm.caching.dual_cache import DualCache + + _, jwk = _get_rsa_key_and_jwk(kid="coalesced-key") + endpoint = _ScriptedJWKSEndpoint(({"keys": [jwk]},), delay=0.05) + jwt_handler = _get_jwt_handler_with_scripted_endpoint(DualCache(), endpoint) + + public_keys = await asyncio.gather( + *[ + jwt_handler._get_public_key_from_jwks_url( + jwks_url="https://coalesce.example.com/keys", + kid="coalesced-key", + ) + for _ in range(5) + ] + ) + + assert public_keys == [jwk] * 5 + assert endpoint.call_count == 1 + + @pytest.mark.asyncio async def test_get_public_key_tries_next_jwks_url_when_kid_missing(monkeypatch): from litellm.caching.dual_cache import DualCache @@ -4140,6 +4710,38 @@ async def test_auth_jwt_issuer_path_expired_token_raises_401(monkeypatch): assert "Token Expired" in exc_info.value.message +@pytest.mark.asyncio +async def test_auth_jwt_issuer_path_unreachable_jwks_raises_503(monkeypatch): + """The issuer-scoped path must report an unreachable IdP as 503, not as a credential failure.""" + monkeypatch.delenv("JWT_AUDIENCE", raising=False) + monkeypatch.delenv("JWT_PUBLIC_KEY_URL", raising=False) + + issuer = "https://unreachable-issuer.example.com" + jwks_url = f"{issuer}/keys" + private_key, _ = _get_rsa_key_and_jwk(kid="unreachable-kid") + + jwt_handler = _get_jwt_handler_with_issuer_keys( + issuers=[{"issuer": issuer, "jwks_url": jwks_url, "audience": "my-audience"}], + keys_by_url={}, + ) + endpoint = _ScriptedJWKSEndpoint((httpx.ConnectTimeout("connect timed out"),)) + jwt_handler.http_handler = endpoint + + token = _encode_rsa_jwt( + private_key=private_key, + issuer=issuer, + audience="my-audience", + kid="unreachable-kid", + ) + + with pytest.raises(ProxyException) as exc_info: + await jwt_handler.auth_jwt(token=token) + + assert exc_info.value.code == "503" + assert exc_info.value.type == ProxyErrorTypes.auth_provider_unavailable + assert endpoint.call_count == 3 + + @pytest.mark.asyncio async def test_multi_issuer_jwt_maps_kubernetes_namespace_claim(monkeypatch): monkeypatch.delenv("JWT_AUDIENCE", raising=False) diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index aa9e5349c87..2f4c2d3870f 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -3428,3 +3428,104 @@ def test_auto_router_dry_runs_share_model_new_audience(user_role, dry_run_route) # Anchor so parity cannot be satisfied by both routes 403ing for everyone if user_role == LitellmUserRoles.INTERNAL_USER.value: assert outcome(dry_run_route) == "allowed" + + +AGENT_MANAGEMENT_ROUTES = [ + "/v1/agents", + "/v1/agents/abc-123", + "/v1/agents/make_public", + "/v1/agents/abc-123/make_public", +] + +AGENT_INFERENCE_ROUTES = [ + "/a2a/abc-123", + "/a2a/abc-123/message/send", + "/a2a/abc-123/message/stream", + "/a2a/abc-123/.well-known/agent-card.json", +] + + +@pytest.mark.parametrize("route", AGENT_MANAGEMENT_ROUTES) +def test_agent_management_routes_classified_as_management_not_llm_api(route): + """Agent registry CRUD must be management routes, not llm_api routes. + + Regression for the Admin UI Agents tab failing with "LLM API routes are + disabled for this instance." on admin nodes that set + DISABLE_LLM_API_ENDPOINTS. + """ + + assert RouteChecks.is_llm_api_route(route=route) is False + assert RouteChecks.is_management_route(route=route) is True + + +@pytest.mark.parametrize("route", AGENT_INFERENCE_ROUTES) +def test_agent_inference_routes_stay_llm_api(route): + """A2A invocation stays on the data plane, gated by DISABLE_LLM_API_ENDPOINTS.""" + + assert RouteChecks.is_llm_api_route(route=route) is True + assert RouteChecks.is_management_route(route=route) is False + + +@pytest.mark.parametrize("route", AGENT_MANAGEMENT_ROUTES + AGENT_INFERENCE_ROUTES) +def test_agent_routes_union_still_covers_both_halves(route): + """Keys configured with allowed_routes=["agent_routes"] must keep both halves.""" + + assert ( + RouteChecks.check_route_access( + route=route, allowed_routes=LiteLLMRoutes.agent_routes.value + ) + is True + ) + + +@pytest.mark.parametrize("route", AGENT_MANAGEMENT_ROUTES) +@pytest.mark.parametrize("method", ["GET", "POST", "DELETE"]) +def test_virtual_key_llm_api_routes_allows_agent_registry(route, method): + """Keys with allowed_routes=["llm_api_routes"] could reach agent CRUD before the + inference/management split and must still reach it after. + + Writes remain proxy-admin-only inside agent_endpoints/endpoints.py, so this + carve-out is not method-aware. + """ + + valid_token = UserAPIKeyAuth(user_id="test_user", allowed_routes=["llm_api_routes"]) + + assert ( + RouteChecks.is_virtual_key_allowed_to_call_route( + route=route, + valid_token=valid_token, + request=_mock_request(method), + ) + is True + ) + + +@pytest.mark.parametrize( + "user_role", + [ + LitellmUserRoles.INTERNAL_USER.value, + LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value, + None, + ], +) +@pytest.mark.parametrize("method, route", [("GET", "/v1/agents"), ("POST", "/v1/agents")]) +def test_agent_registry_route_gate_open_to_non_admin_roles(user_role, method, route): + """Non-admin callers reached agent CRUD through llm_api_routes before the split. + + The route gate must keep letting them through so the handlers can scope the + listing by role and 403 non-admin writes themselves. + """ + + valid_token = UserAPIKeyAuth(user_id="test_user", user_role=user_role) + request = MagicMock(spec=Request) + request.method = method + request.query_params = {} + + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=LiteLLM_UserTable(user_id="test_user", user_role=user_role), + _user_role=user_role, + route=route, + request=request, + valid_token=valid_token, + request_data={}, + ) diff --git a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py index 608dc8cb5c8..c5bc4e29f81 100644 --- a/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py +++ b/tests/test_litellm/proxy/common_utils/test_reset_budget_job.py @@ -8,6 +8,8 @@ from datetime import time as dt_time from typing import Any, Dict, List from unittest.mock import AsyncMock, MagicMock +import httpx +import prisma import pytest sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system path @@ -1932,3 +1934,311 @@ def test_user_and_team_chunks_report_progress_despite_a_failed_row( assert client.fetches_by_table[table_name] == 2 assert len(_batch_writes(client, table_name, op="update")) == 2 assert [call["call_type"] for call in logging_obj.service_logging_obj.failure_calls] == [call_type] + + +class FlakyPrismaClient(MockPrismaClient): + """A client whose first N reads (or first N batch commits) fail with a + transport error, and which records every reconnect attempt. + """ + + def __init__(self, *, read_failures: int = 0, commit_failures: int = 0, error: Exception | None = None): + super().__init__() + self.reconnect_reasons: List[str] = [] + self.read_attempts: int = 0 + self.commit_attempts: int = 0 + self._read_failures = read_failures + self._commit_failures = commit_failures + self._error = error or httpx.ConnectError("All connection attempts failed") + + outer = self + original_batch = self.db.batch_ + + def _batch_(): + batcher = original_batch() + batch_commit = batcher.commit + + async def _maybe_failing_commit(): + outer.commit_attempts += 1 + if outer._commit_failures > 0: + outer._commit_failures -= 1 + raise outer._error + return await batch_commit() + + batcher.commit = _maybe_failing_commit + return batcher + + self.db.batch_ = _batch_ + + async def attempt_db_reconnect(self, *, reason, timeout_seconds=None, lock_timeout_seconds=None) -> bool: + self.reconnect_reasons.append(reason) + return True + + async def get_data(self, table_name, query_type, **kwargs): + self.read_attempts += 1 + if self._read_failures > 0: + self._read_failures -= 1 + raise self._error + return await super().get_data(table_name, query_type, **kwargs) + + +def _due_row(table: str, identifier: str): + now = datetime.now(timezone.utc) + id_field = {"key": "token", "user": "user_id", "team": "team_id"}[table] + return type( + "Row", + (), + { + "spend": _DUE_ROW_SPEND, + "budget_duration": "30d", + "budget_reset_at": now - timedelta(seconds=1), + id_field: identifier, + }, + ) + + +@pytest.mark.parametrize( + "phase, table_name, reason", + [ + ("reset_budget_for_litellm_keys", "key", "reset_budget_read_keys_failure"), + ("reset_budget_for_litellm_users", "user", "reset_budget_read_users_failure"), + ("reset_budget_for_litellm_teams", "team", "reset_budget_read_teams_failure"), + ], + ids=["keys", "users", "teams"], +) +def test_transient_transport_error_on_read_reconnects_and_still_resets(phase, table_name, reason): + """A dropped connection on the read must cost one reconnect-and-retry, not + the whole tick (LIT-5372). Pre-fix the httpx.ConnectError was swallowed and + the phase reset nothing until the next tick, 10 minutes later. + """ + client = FlakyPrismaClient(read_failures=1) + client.data[table_name] = [_due_row(table_name, "row-1")] + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + asyncio.run(getattr(job, phase)()) + + assert client.reconnect_reasons == [reason] + assert len(_batch_writes(client, table_name, op="update")) == 1 + + +@pytest.mark.parametrize( + "phase, table_name, reason", + [ + ("reset_budget_for_litellm_keys", "key", "reset_budget_write_keys_failure"), + ("reset_budget_for_litellm_users", "user", "reset_budget_write_users_failure"), + ("reset_budget_for_litellm_teams", "team", "reset_budget_write_teams_failure"), + ], + ids=["keys", "users", "teams"], +) +def test_connect_error_on_write_reconnects_and_commits(phase, table_name, reason): + """A ConnectError proves the commit never reached the database, so replaying + it cannot double-apply anything: the rows still get reset on this tick. + """ + client = FlakyPrismaClient(commit_failures=1) + client.data[table_name] = [_due_row(table_name, "row-1")] + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + asyncio.run(getattr(job, phase)()) + + assert client.reconnect_reasons == [reason] + assert client.commit_attempts == 2 + assert len(_batch_writes(client, table_name, op="update")) == 1 + + +@pytest.mark.parametrize("ambiguous_error_name", ["ReadError", "ReadTimeout"]) +def test_ambiguous_transport_error_on_write_is_not_replayed(ambiguous_error_name): + """A post-send transport error leaves the commit outcome unknown. Since the + reset zeroes spend unconditionally, replaying it would erase spend accrued + after a commit that actually landed, so only reads may retry these. + """ + client = FlakyPrismaClient( + commit_failures=1, + error=getattr(httpx, ambiguous_error_name)("ambiguous"), + ) + client.data["key"] = [_due_row("key", "tok-1")] + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.reconnect_reasons == [] + assert client.commit_attempts == 1 + + +@pytest.mark.parametrize("ambiguous_error_name", ["ReadError", "ReadTimeout"]) +def test_ambiguous_transport_error_on_read_still_retries(ambiguous_error_name): + """Reads have nothing to double-apply, so the full transport class retries.""" + client = FlakyPrismaClient(read_failures=1, error=getattr(httpx, ambiguous_error_name)("ambiguous")) + client.data["key"] = [_due_row("key", "tok-1")] + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.reconnect_reasons == ["reset_budget_read_keys_failure"] + assert len(_batch_writes(client, "key", op="update")) == 1 + + +def test_transport_error_on_budget_cascade_read_reconnects_and_commits(): + client = FlakyPrismaClient(read_failures=1) + budget = _budget_row(budget_id="b-1", budget_duration="1d") + client.data["budget"] = [budget] + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + asyncio.run(job.reset_budget_for_litellm_budget_table()) + + assert client.reconnect_reasons == ["reset_budget_read_budgets_failure"] + assert [w["where"]["budget_id"] for w in _batch_writes(client, "budget", op="update_many")] == ["b-1"] + + +def test_non_transport_error_still_surfaces_without_a_reconnect(): + """A UniqueViolationError means the DB is reachable and the statement was + refused, so reconnecting would be pointless: the phase must fail as before. + """ + client = FlakyPrismaClient(read_failures=1, error=prisma.errors.UniqueViolationError(MagicMock())) + client.data["key"] = [_due_row("key", "tok-1")] + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.reconnect_reasons == [] + assert client.read_attempts == 1 + assert _batch_writes(client, "key") == [] + + +def test_transport_error_that_outlives_the_reconnect_is_not_retried_forever(): + client = FlakyPrismaClient(read_failures=2) + client.data["key"] = [_due_row("key", "tok-1")] + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.reconnect_reasons == ["reset_budget_read_keys_failure"] + assert client.read_attempts == 2 + assert _batch_writes(client, "key") == [] + + +def test_transport_error_on_window_read_reconnects_and_still_resets(monkeypatch): + """The raw per-window queries are reads too, so a blip there must not cost + the whole window-reset phase.""" + expired = (datetime.utcnow() - timedelta(minutes=5)).isoformat() + "Z" + key_rows = [{"token": "sk-expired", "budget_limits": [{"budget_duration": "1d", "reset_at": expired}]}] + job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=key_rows, team_rows=[]) + reconnect_reasons: List[str] = [] + good_query_raw = prisma_client.db.query_raw + + async def failing_once_query_raw(query: str, *args, **kwargs): + if '"LiteLLM_VerificationToken"' in query and not reconnect_reasons: + raise httpx.ConnectError("All connection attempts failed") + return await good_query_raw(query, *args, **kwargs) + + async def record_reconnect(*, reason, timeout_seconds=None, lock_timeout_seconds=None) -> bool: + reconnect_reasons.append(reason) + return True + + prisma_client.db.query_raw = AsyncMock(side_effect=failing_once_query_raw) + prisma_client.attempt_db_reconnect = record_reconnect + + asyncio.run(job.reset_budget_windows()) + + assert reconnect_reasons == ["reset_budget_read_key_windows_failure"] + prisma_client.db.litellm_verificationtoken.update.assert_awaited_once() + + +def test_connect_error_on_window_write_reconnects_and_writes(monkeypatch): + expired = (datetime.utcnow() - timedelta(minutes=5)).isoformat() + "Z" + team_rows = [{"team_id": "team-expired", "budget_limits": [{"budget_duration": "1d", "reset_at": expired}]}] + job, prisma_client, _ = _make_reset_budget_windows_job(monkeypatch, key_rows=[], team_rows=team_rows) + reconnect_reasons: List[str] = [] + + async def failing_once_update(**kwargs) -> None: + if not reconnect_reasons: + raise httpx.ConnectError("All connection attempts failed") + + async def record_reconnect(*, reason, timeout_seconds=None, lock_timeout_seconds=None) -> bool: + reconnect_reasons.append(reason) + return True + + prisma_client.db.litellm_teamtable.update = AsyncMock(side_effect=failing_once_update) + prisma_client.attempt_db_reconnect = record_reconnect + + asyncio.run(job.reset_budget_windows()) + + assert reconnect_reasons == ["reset_budget_write_team_windows_failure"] + assert prisma_client.db.litellm_teamtable.update.await_count == 2 + + +_DUE_ROW_SPEND = 42.0 +_SPEND_ACCRUED_AFTER_COMMIT = 7.5 + + +class AmbiguousCommitClient(MockPrismaClient): + """A client whose batch commit lands in the database and only then fails in + transit, so the caller cannot tell whether it committed. + + The queued spend-zero is applied to `key_spend`, and fresh usage accrues in + the window between that landed commit and any replay, so a replay is + observable as erased spend rather than merely as an extra commit. + """ + + def __init__(self, *, error: Exception, spend_accrued_after_commit: float): + super().__init__() + self.key_spend: float = _DUE_ROW_SPEND + self.commit_attempts: int = 0 + self.reconnect_reasons: list[str] = [] + + outer = self + original_batch = self.db.batch_ + + def _batch_(): + batcher = original_batch() + batch_commit = batcher.commit + + async def _commit_then_lose_the_response(): + outer.commit_attempts += 1 + result = await batch_commit() + for call in batcher.calls: + if call["table"] == "key" and call["data"].get("spend") == 0: + outer.key_spend = 0.0 + if outer.commit_attempts > 1: + return result + outer.key_spend += spend_accrued_after_commit + raise error + + batcher.commit = _commit_then_lose_the_response + return batcher + + self.db.batch_ = _batch_ + + async def attempt_db_reconnect(self, *, reason, timeout_seconds=None, lock_timeout_seconds=None) -> bool: + self.reconnect_reasons.append(reason) + return True + + +@pytest.mark.parametrize( + "error, expected_commits, expected_spend, expected_reconnects", + [ + (httpx.ReadError("response lost in transit"), 1, _SPEND_ACCRUED_AFTER_COMMIT, []), + (httpx.ReadTimeout("response lost in transit"), 1, _SPEND_ACCRUED_AFTER_COMMIT, []), + (httpx.ConnectError("never left the client"), 2, 0.0, ["reset_budget_write_keys_failure"]), + ], + ids=["read_error", "read_timeout", "connect_error_erasure_control"], +) +def test_ambiguous_commit_replay_does_not_erase_newly_accrued_spend( + error, expected_commits, expected_spend, expected_reconnects +): + """A reset zeroes spend unconditionally, so replaying a commit that already + landed erases every dollar spent since it landed (LIT-5372 review finding). + + The `connect_error` case is the control: it is the one error class allowed + to replay, and driving it through this same land-then-fail harness proves + the spend assertion can actually observe an erasure. In production a + ConnectError means the statements never reached the database, so its replay + has nothing to erase. + """ + client = AmbiguousCommitClient(error=error, spend_accrued_after_commit=_SPEND_ACCRUED_AFTER_COMMIT) + client.data["key"] = [_due_row("key", "tok-1")] + job = ResetBudgetJob(proxy_logging_obj=MockProxyLogging(), prisma_client=client) + + asyncio.run(job.reset_budget_for_litellm_keys()) + + assert client.key_spend == expected_spend + assert client.commit_attempts == expected_commits + assert client.reconnect_reasons == expected_reconnects diff --git a/tests/test_litellm/proxy/db/conftest.py b/tests/test_litellm/proxy/db/conftest.py index a0fb6bed4fa..6f67b91ac1d 100644 --- a/tests/test_litellm/proxy/db/conftest.py +++ b/tests/test_litellm/proxy/db/conftest.py @@ -6,6 +6,7 @@ import pytest DB_ENV_KEYS = ( "IAM_TOKEN_DB_AUTH", + "AZURE_POSTGRESQL_AUTH", "DATABASE_URL", "DIRECT_URL", "DATABASE_URL_READ_REPLICA", @@ -59,6 +60,17 @@ def pytest_runtest_teardown(item: pytest.Item, nextitem: Optional[pytest.Item]) return result +@pytest.fixture(autouse=True) +def reset_entra_token_provider_cache() -> Generator[None, None, None]: + """The Entra provider factory is cached process-wide so one Azure credential serves + the whole proxy; that cache would otherwise carry one test's stub into the next.""" + from litellm.proxy.db.token_auth import build_azure_entra_token_provider + + build_azure_entra_token_provider.cache_clear() + yield + build_azure_entra_token_provider.cache_clear() + + @pytest.fixture def unset_database_url(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setenv("DATABASE_URL", "about-to-be-unset") diff --git a/tests/test_litellm/proxy/db/mcp_server/test_db.py b/tests/test_litellm/proxy/db/mcp_server/test_db.py index ff6400ac5b7..481d1a864c0 100644 --- a/tests/test_litellm/proxy/db/mcp_server/test_db.py +++ b/tests/test_litellm/proxy/db/mcp_server/test_db.py @@ -1,6 +1,7 @@ -import json import os import sys +from types import SimpleNamespace +from unittest.mock import AsyncMock, MagicMock import pytest @@ -11,5 +12,34 @@ sys.path.insert( from litellm.proxy._experimental.mcp_server.db import get_mcp_servers_by_team -def test_fetch_mcp_servers_by_team(): - assert True == True +def _prisma_client_returning(team_record: object) -> MagicMock: + prisma_client = MagicMock() + prisma_client.db.litellm_teamtable.find_unique = AsyncMock(return_value=team_record) + return prisma_client + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "team_record, expected", + [ + (None, []), + (SimpleNamespace(object_permission=None), []), + (SimpleNamespace(object_permission=SimpleNamespace(mcp_servers=None)), []), + (SimpleNamespace(object_permission=SimpleNamespace(mcp_servers=[])), []), + ( + SimpleNamespace( + object_permission=SimpleNamespace(mcp_servers=["server_a", "server_b"]) + ), + ["server_a", "server_b"], + ), + ], +) +async def test_fetch_mcp_servers_by_team(team_record, expected): + prisma_client = _prisma_client_returning(team_record) + + assert await get_mcp_servers_by_team(prisma_client, "team-123") == expected + + prisma_client.db.litellm_teamtable.find_unique.assert_awaited_once_with( + where={"team_id": "team-123"}, + include={"object_permission": True}, + ) diff --git a/tests/test_litellm/proxy/db/test_db_url_settings.py b/tests/test_litellm/proxy/db/test_db_url_settings.py index e5aa09addab..e83e8310626 100644 --- a/tests/test_litellm/proxy/db/test_db_url_settings.py +++ b/tests/test_litellm/proxy/db/test_db_url_settings.py @@ -3,8 +3,8 @@ The model assembles ``DATABASE_URL`` (and optionally ``DATABASE_URL_READ_REPLICA``) from the discrete ``DATABASE_*`` env vars emitted by the ``helm/litellm`` chart, before Prisma initializes. It covers -both IAM auth (mint a short-lived token) and password auth, for both the -writer and the read replica. +both token auth (mint a short-lived AWS RDS IAM or Microsoft Entra ID token) +and password auth, for both the writer and the read replica. The reader URL is opt-in via ``DATABASE_HOST_READ_REPLICA`` and must not clobber a pre-existing ``DATABASE_URL_READ_REPLICA``. A pre-existing @@ -15,12 +15,14 @@ import os from unittest.mock import patch import pytest +from pydantic import ValidationError from litellm.proxy.db.db_url_settings import ( DatabaseURLSettings, unsupported_db_scheme, unsupported_db_scheme_message, ) +from litellm.proxy.db.token_auth import AzureEntraTokenAuth, RdsIamTokenAuth def _apply() -> bool: @@ -30,6 +32,7 @@ def _apply() -> bool: _MANAGED_DB_ENV_VARS = ( "IAM_TOKEN_DB_AUTH", + "AZURE_POSTGRESQL_AUTH", "DATABASE_URL", "DIRECT_URL", "DATABASE_URL_READ_REPLICA", @@ -76,6 +79,14 @@ def _stub_iam_token(token: str = "FAKE_TOKEN"): ) +def _stub_entra_token(token: str = "FAKE_TOKEN"): + """Patch the Azure-touching token provider so tests don't need azure-identity.""" + return patch( + "litellm.secret_managers.get_azure_ad_token_provider.get_azure_ad_token_provider", + return_value=lambda: token, + ) + + # --------------------------------------------------------------------------- # IAM auth # --------------------------------------------------------------------------- @@ -104,6 +115,35 @@ def test_assembles_writer_url_when_iam_enabled(monkeypatch): assert "DATABASE_URL_READ_REPLICA" not in os.environ +def test_a_pre_encoded_iam_user_survives_url_assembly(monkeypatch): + """This URL used to be interpolated raw, so pre-encoding ``DATABASE_USER`` was the + only way to run IAM auth as a user whose name contains an ``@``. Encoding it again + yields ``svc%2540corp``, which Postgres rejects with + ``User `svc%40corp` was denied access``.""" + monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true") + monkeypatch.setenv("DATABASE_HOST", "writer.example.com") + monkeypatch.setenv("DATABASE_USER", "svc%40corp") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + + with _stub_iam_token("WRITER_TOKEN"): + assert _apply() is True + + assert os.environ["DATABASE_URL"] == "postgresql://svc%40corp:WRITER_TOKEN@writer.example.com:5432/litellm_db" + + +def test_an_unreadable_toggle_fails_the_settings_model(monkeypatch): + """Pydantic rejected `IAM_TOKEN_DB_AUTH=enabled` before token auth had its own + parser. Reading it as 'off' instead would silently drop an operator who asked for + token auth down to password auth, with no log line saying so.""" + monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "enabled") + monkeypatch.setenv("DATABASE_HOST", "writer.example.com") + monkeypatch.setenv("DATABASE_USER", "litellm") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + + with pytest.raises(ValidationError, match="IAM_TOKEN_DB_AUTH"): + DatabaseURLSettings.from_env() + + def test_missing_writer_envs_raises(monkeypatch): monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true") # DATABASE_HOST intentionally unset. @@ -184,6 +224,123 @@ def test_reader_field_fallbacks_default_to_writer_values(monkeypatch): ) +# --------------------------------------------------------------------------- +# Azure Entra auth +# --------------------------------------------------------------------------- + + +def test_assembles_writer_url_when_azure_entra_enabled(monkeypatch): + monkeypatch.setenv("AZURE_POSTGRESQL_AUTH", "true") + monkeypatch.setenv("DATABASE_HOST", "writer.postgres.database.azure.com") + monkeypatch.setenv("DATABASE_USER", "litellm@contoso.onmicrosoft.com") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + + with _stub_entra_token("ENTRA_TOKEN"): + assert _apply() is True + + assert os.environ["DATABASE_URL"] == ( + "postgresql://litellm%40contoso.onmicrosoft.com:ENTRA_TOKEN" + "@writer.postgres.database.azure.com:5432/litellm_db" + ) + assert os.environ["AZURE_POSTGRESQL_AUTH"] == "True" + assert "IAM_TOKEN_DB_AUTH" not in os.environ + + +def test_azure_reader_url_assembled_from_writer_fallbacks(monkeypatch): + monkeypatch.setenv("AZURE_POSTGRESQL_AUTH", "true") + monkeypatch.setenv("DATABASE_HOST", "writer.postgres.database.azure.com") + monkeypatch.setenv("DATABASE_USER", "litellm@contoso.onmicrosoft.com") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + monkeypatch.setenv("DATABASE_SCHEMA", "public") + monkeypatch.setenv("DATABASE_HOST_READ_REPLICA", "reader.postgres.database.azure.com") + + with _stub_entra_token("ENTRA_TOKEN"): + _apply() + + assert os.environ["DATABASE_URL_READ_REPLICA"] == ( + "postgresql://litellm%40contoso.onmicrosoft.com:ENTRA_TOKEN" + "@reader.postgres.database.azure.com:5432/litellm_db?schema=public" + ) + + +def test_azure_missing_writer_envs_names_the_azure_toggle(monkeypatch): + monkeypatch.setenv("AZURE_POSTGRESQL_AUTH", "true") + # DATABASE_HOST intentionally unset. + monkeypatch.setenv("DATABASE_USER", "litellm@contoso.onmicrosoft.com") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + + with pytest.raises(RuntimeError, match="AZURE_POSTGRESQL_AUTH is enabled but"): + _apply() + + +def test_both_token_toggles_is_a_startup_error(monkeypatch): + monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "true") + monkeypatch.setenv("AZURE_POSTGRESQL_AUTH", "true") + monkeypatch.setenv("DATABASE_HOST", "writer.example.com") + monkeypatch.setenv("DATABASE_USER", "litellm") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + + with pytest.raises(RuntimeError, match="can only come from one token source"): + _apply() + + assert "DATABASE_URL" not in os.environ + + +@pytest.mark.parametrize( + "env_var, expected_type", + [("IAM_TOKEN_DB_AUTH", RdsIamTokenAuth), ("AZURE_POSTGRESQL_AUTH", AzureEntraTokenAuth)], +) +def test_token_auth_reflects_the_enabled_toggle(monkeypatch, env_var, expected_type): + monkeypatch.setenv(env_var, "true") + + with _stub_entra_token(): + assert isinstance(DatabaseURLSettings.from_env().token_auth(), expected_type) + + +def test_the_toggle_agrees_with_the_refresh_loop_on_every_spelling(monkeypatch): + """This model and `resolve_database_token_auth` (which arms the refresh loop) both + read the same env var. When they disagreed, `AZURE_POSTGRESQL_AUTH=1` minted a token + here and left the refresh loop convinced token auth was off.""" + from litellm.proxy.db.token_auth import resolve_database_token_auth + + monkeypatch.setenv("AZURE_POSTGRESQL_AUTH", "1") + + with _stub_entra_token(): + settings_says = DatabaseURLSettings.from_env().azure_postgresql_auth + refresh_loop_says = resolve_database_token_auth() is not None + + assert settings_says is True + assert refresh_loop_says is True + + +def test_an_empty_toggle_is_off_rather_than_a_validation_error(monkeypatch): + """`value: ""` is how a Kubernetes manifest spells 'off', and the componentized + entrypoints build this model at import time, so a raise there is a crash loop.""" + monkeypatch.setenv("AZURE_POSTGRESQL_AUTH", "") + monkeypatch.setenv("IAM_TOKEN_DB_AUTH", "") + + settings = DatabaseURLSettings.from_env() + + assert (settings.azure_postgresql_auth, settings.iam_token_db_auth) == (False, False) + assert settings.token_auth() is None + + +def test_apply_writer_url_to_env_leaves_the_reader_alone(monkeypatch): + """The CLI shares the writer minting path but resolves the read replica itself, so + it must not start writing DATABASE_URL_READ_REPLICA as a side effect.""" + monkeypatch.setenv("AZURE_POSTGRESQL_AUTH", "true") + monkeypatch.setenv("DATABASE_HOST", "writer.postgres.database.azure.com") + monkeypatch.setenv("DATABASE_USER", "litellm") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + monkeypatch.setenv("DATABASE_HOST_READ_REPLICA", "reader.postgres.database.azure.com") + + with _stub_entra_token("ENTRA_TOKEN"): + assert DatabaseURLSettings.from_env().apply_writer_url_to_env() is True + + assert "DATABASE_URL" in os.environ + assert "DATABASE_URL_READ_REPLICA" not in os.environ + + # --------------------------------------------------------------------------- # Password auth # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/db/test_exception_handler.py b/tests/test_litellm/proxy/db/test_exception_handler.py index 474e571e592..84a9ddfacff 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler.py +++ b/tests/test_litellm/proxy/db/test_exception_handler.py @@ -117,6 +117,8 @@ def test_is_database_connection_generic_errors(): TimeoutError("timed out"), OSError("network is unreachable"), asyncio.TimeoutError(), + httpx.ConnectError("connection refused"), + httpx.ConnectTimeout("connect timed out"), HTTPClientClosedError(), ClientNotConnectedError(), PrismaError("can't reach database server"), diff --git a/tests/test_litellm/proxy/db/test_exception_handler_reconnect_retry.py b/tests/test_litellm/proxy/db/test_exception_handler_reconnect_retry.py index ae0e1f845b0..0a25ed55e90 100644 --- a/tests/test_litellm/proxy/db/test_exception_handler_reconnect_retry.py +++ b/tests/test_litellm/proxy/db/test_exception_handler_reconnect_retry.py @@ -14,7 +14,7 @@ from unittest.mock import AsyncMock, MagicMock import httpx import pytest -from prisma.errors import UniqueViolationError +from prisma.errors import ClientNotConnectedError, UniqueViolationError sys.path.insert(0, os.path.abspath("../../..")) @@ -253,3 +253,45 @@ async def test_call_with_db_reconnect_retry_preserves_original_error_when_reconn assert exc_info.value is original_exc assert exc_info.value.__cause__ is reconnect_exc client.attempt_db_reconnect.assert_awaited_once() + +@pytest.mark.asyncio +async def test_call_with_db_reconnect_retry_honors_narrowed_retry_safe_types(): + """A non-idempotent write can pass `retry_safe_error_types` to opt out of + replaying post-send transport errors, whose commit outcome is unknown.""" + client = _make_client(attempt_db_reconnect_return=True) + attempts = 0 + + async def _factory(): + nonlocal attempts # rebind-ok: attempt counter for a two-call helper + attempts += 1 + raise httpx.ReadError("ambiguous") + + with pytest.raises(httpx.ReadError): + await call_with_db_reconnect_retry( + client, + _factory, + reason="write_narrowed", + retry_safe_error_types=(httpx.ConnectError,), + ) + + assert attempts == 1 + client.attempt_db_reconnect.assert_not_awaited() + + +@pytest.mark.asyncio +async def test_call_with_db_reconnect_retry_default_covers_every_transport_error(): + """Callers that don't narrow keep retrying anything + `is_database_transport_error` accepts, not just the httpx types.""" + client = _make_client(attempt_db_reconnect_return=True) + attempts = 0 + + async def _factory(): + nonlocal attempts # rebind-ok: attempt counter for a two-call helper + attempts += 1 + if attempts == 1: + raise ClientNotConnectedError() + return "ok" + + assert await call_with_db_reconnect_retry(client, _factory, reason="default_wide") == "ok" + assert attempts == 2 + client.attempt_db_reconnect.assert_awaited_once() diff --git a/tests/test_litellm/proxy/db/test_prisma_client.py b/tests/test_litellm/proxy/db/test_prisma_client.py index 08b873dfc44..395f17e85ef 100644 --- a/tests/test_litellm/proxy/db/test_prisma_client.py +++ b/tests/test_litellm/proxy/db/test_prisma_client.py @@ -2,6 +2,7 @@ import json import os import signal import sys +import urllib.parse from unittest.mock import AsyncMock, MagicMock, Mock, patch import pytest @@ -215,3 +216,129 @@ def test_db_push_applies_replica_identity_full_when_requested(monkeypatch): assert mock_run.call_args[0][0][:3] == ["prisma", "db", "push"] assert applied == [True] + + +def _entra_jwt(expires_in_seconds: int) -> str: + """A JWT shaped like a real Entra access token, expiring ``expires_in_seconds`` from now.""" + import base64 + from datetime import datetime, timedelta, timezone + + exp = int((datetime.now(tz=timezone.utc) + timedelta(seconds=expires_in_seconds)).timestamp()) + payload = base64.urlsafe_b64encode(json.dumps({"exp": exp}).encode()).rstrip(b"=").decode() + return f"aGVhZGVy.{payload}.c2ln" + + +@pytest.fixture +def azure_env(monkeypatch, unset_database_url): + monkeypatch.setenv("DATABASE_HOST", "pg.postgres.database.azure.com") + monkeypatch.setenv("DATABASE_PORT", "5432") + monkeypatch.setenv("DATABASE_USER", "litellm@contoso.onmicrosoft.com") + monkeypatch.setenv("DATABASE_NAME", "litellm_db") + + +def _azure_wrapper(token: str, **kwargs): + from litellm.proxy.db.token_auth import AzureEntraTokenAuth + + return PrismaWrapper( + original_prisma=MagicMock(), + token_auth=AzureEntraTokenAuth(token_provider=lambda: token), + **kwargs, + ) + + +def test_azure_entra_mint_writes_an_encoded_url_into_the_db_url_env_var(azure_env): + """The UPN user and the JWT both have to survive being embedded in a URL.""" + token = _entra_jwt(3600) + wrapper = _azure_wrapper(token) + + db_url = wrapper.get_rds_iam_token() + + assert db_url == ( + f"postgresql://litellm%40contoso.onmicrosoft.com:{urllib.parse.quote(token, safe='')}" + "@pg.postgres.database.azure.com:5432/litellm_db" + ) + assert os.environ["DATABASE_URL"] == db_url + + +def test_azure_entra_refresh_is_scheduled_off_the_jwt_expiry(azure_env): + """Without reading `exp` this falls back to a fixed 600s interval, which silently + outlives a token and breaks every reconnect after it lapses (issue #29661).""" + wrapper = _azure_wrapper(_entra_jwt(3600)) + wrapper.get_rds_iam_token() + + seconds = wrapper._calculate_seconds_until_refresh() + + expected = 3600 - PrismaWrapper.TOKEN_REFRESH_BUFFER_SECONDS + assert seconds != PrismaWrapper.FALLBACK_REFRESH_INTERVAL_SECONDS + assert expected - 5 <= seconds <= expected + + +def test_a_token_whose_expiry_never_advances_cannot_spin_the_refresh_loop(azure_env): + """azure-identity hands back its cached token when a renewal attempt fails inside its + own window, so a transient Entra or IMDS problem in the last 3 minutes of a token + yields a successful refresh whose `exp` has not moved. With no floor on the sleep the + loop then re-mints and recreates the query engine on every pass, with nothing in + between, for as long as Entra stays sick.""" + wrapper = _azure_wrapper(_entra_jwt(60)) + wrapper.get_rds_iam_token() + first = wrapper._calculate_seconds_until_refresh() + + wrapper.get_rds_iam_token() + second = wrapper._calculate_seconds_until_refresh() + + assert first == second == PrismaWrapper.TOKEN_REFRESH_MIN_SLEEP_SECONDS + + +def test_azure_entra_token_expiry_is_detected(azure_env): + wrapper = _azure_wrapper(_entra_jwt(3600)) + fresh_url = wrapper.get_rds_iam_token() + expired_url = _azure_wrapper(_entra_jwt(-1)).get_rds_iam_token() + + assert wrapper.is_token_expired(fresh_url) is False + assert wrapper.is_token_expired(expired_url) is True + + +@pytest.mark.asyncio +async def test_azure_entra_strategy_starts_the_refresh_task(azure_env): + """The refresh loop is gated on the legacy boolean, so an Azure strategy has to + get past that gate; a password-auth wrapper still must not start a task.""" + wrapper = _azure_wrapper(_entra_jwt(3600)) + wrapper.get_rds_iam_token() + password_wrapper = PrismaWrapper(original_prisma=MagicMock()) + + await wrapper.start_token_refresh_task() + await password_wrapper.start_token_refresh_task() + try: + assert wrapper._token_refresh_task is not None + assert not wrapper._token_refresh_task.done() + assert password_wrapper._token_refresh_task is None + finally: + await wrapper.stop_token_refresh_task() + + +def test_azure_entra_strategy_reads_as_token_auth_enabled(azure_env): + """`routing_prisma_wrapper` gates the reader's refresh on this flag, so an Azure + reader has to answer True to it.""" + wrapper = _azure_wrapper(_entra_jwt(3600)) + + assert wrapper.iam_token_db_auth is True + assert wrapper.token_label == "Azure Entra token" + + +def test_the_token_strategy_cannot_be_swapped_after_construction(azure_env): + """Assigning the legacy boolean used to replace a configured Entra strategy with the + RDS one, which points boto at an Azure host.""" + wrapper = _azure_wrapper(_entra_jwt(3600)) + + with pytest.raises(AttributeError): + wrapper.iam_token_db_auth = True + + +def test_minting_without_the_database_env_vars_names_them(azure_env, monkeypatch): + """A blank host used to produce `postgresql://:@:5432/`, which fails deep + inside Prisma instead of at the misconfiguration.""" + monkeypatch.delenv("DATABASE_HOST") + wrapper = _azure_wrapper(_entra_jwt(3600)) + + with pytest.raises(RuntimeError, match="DATABASE_HOST"): + wrapper.get_rds_iam_token() diff --git a/tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py b/tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py index e5bb8b99507..11ed63cf8f0 100644 --- a/tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py +++ b/tests/test_litellm/proxy/db/test_routing_prisma_wrapper.py @@ -991,3 +991,52 @@ async def test_recreate_keeps_writer_unavailable_when_writer_recreate_fails(): await routing.recreate_prisma_client("writer-url") assert routing.writer_unavailable is True + + +def test_prisma_client_premints_an_entra_token_for_the_reader(monkeypatch): + """Under Azure Entra auth the reader has to be pre-minted the same way the RDS + reader already is: Prisma is constructed with a `datasource` URL, so a reader built + from the operator's placeholder URL would never carry a real token.""" + from litellm.proxy.db.prisma_client import PrismaWrapper + from litellm.proxy.db.routing_prisma_wrapper import RoutingPrismaWrapper + from litellm.proxy.db.token_auth import AzureEntraTokenAuth + + monkeypatch.setenv("AZURE_POSTGRESQL_AUTH", "true") + monkeypatch.delenv("IAM_TOKEN_DB_AUTH", raising=False) + monkeypatch.setenv( + "DATABASE_URL_READ_REPLICA", + "postgresql://litellm%40contoso.com@reader.postgres.database.azure.com:5432/litellm", + ) + + captured_kwargs: Dict[str, Any] = {} + + class FakePrisma: + def __init__(self, **kwargs): + captured_kwargs.update(kwargs) + + async def connect(self): + return None + + fake_prisma_module = MagicMock() + fake_prisma_module.Prisma = FakePrisma + monkeypatch.setitem(sys.modules, "prisma", fake_prisma_module) + + with patch( + "litellm.secret_managers.get_azure_ad_token_provider.get_azure_ad_token_provider", + return_value=lambda: "ENTRA-TOKEN", + ): + from litellm.proxy.utils import PrismaClient + + client = PrismaClient( + database_url="postgresql://litellm@writer.postgres.database.azure.com:5432/litellm", + proxy_logging_obj=MagicMock(), + ) + + assert isinstance(client.db, RoutingPrismaWrapper) + assert captured_kwargs["datasource"] == { + "url": "postgresql://litellm%40contoso.com:ENTRA-TOKEN@reader.postgres.database.azure.com:5432/litellm" + } + assert os.environ["DATABASE_URL_READ_REPLICA"] == captured_kwargs["datasource"]["url"] + assert isinstance(client.db._reader.token_auth, AzureEntraTokenAuth) + assert isinstance(client.db._writer.token_auth, AzureEntraTokenAuth) + assert isinstance(client.db._writer, PrismaWrapper) diff --git a/tests/test_litellm/proxy/db/test_token_auth.py b/tests/test_litellm/proxy/db/test_token_auth.py new file mode 100644 index 00000000000..56bdcd6f3e3 --- /dev/null +++ b/tests/test_litellm/proxy/db/test_token_auth.py @@ -0,0 +1,329 @@ +"""Tests for the database token auth strategies. + +``litellm/proxy/db/token_auth.py`` decides where the proxy's Postgres password +comes from: an AWS RDS IAM token, a Microsoft Entra ID access token for Azure +Database for PostgreSQL, or neither. Minting and expiry parsing dispatch over +that union, so both variants are exercised here, together with the URL encoding +that lets an Entra principal (a UPN containing ``@``) survive being embedded in +a connection URL. +""" + +import base64 +import json +from datetime import datetime, timezone +from unittest.mock import patch + +import pytest + +from litellm.proxy.db.token_auth import ( + AZURE_POSTGRESQL_AUTH_ENV_VAR, + AZURE_POSTGRESQL_SCOPE, + IAM_TOKEN_DB_AUTH_ENV_VAR, + AzureEntraTokenAuth, + IAMEndpoint, + RdsIamTokenAuth, + build_azure_entra_token_provider, + mint_database_token, + parse_database_token_expiration, + parse_iam_endpoint_from_url, + resolve_database_token_auth, +) + + +def _entra_token(exp: int, *, header: str = "eyJhbGciOiJSUzI1NiJ9") -> str: + """A JWT shaped like a real Entra access token, carrying ``exp``.""" + payload = base64.urlsafe_b64encode( + json.dumps({"aud": "https://ossrdbms-aad.database.windows.net", "exp": exp}).encode() + ).rstrip(b"=") + return f"{header}.{payload.decode()}.c2lnbmF0dXJl" + + +def _endpoint(**overrides) -> IAMEndpoint: + fields = { + "host": "pg.postgres.database.azure.com", + "port": "5432", + "user": "litellm", + "name": "litellm_db", + } + fields.update(overrides) + return IAMEndpoint(**fields) + + +# --------------------------------------------------------------------------- +# Minting +# --------------------------------------------------------------------------- + + +def test_rds_mint_delegates_to_the_sigv4_token_generator(): + endpoint = _endpoint(host="writer.aurora.local", user="litellm_rds") + + with patch( + "litellm.proxy.auth.rds_iam_token.generate_iam_auth_token", + return_value="SIGV4_TOKEN", + ) as generate: + token = mint_database_token(RdsIamTokenAuth(), endpoint) + + assert token == "SIGV4_TOKEN" + generate.assert_called_once_with( + db_host="writer.aurora.local", + db_port="5432", + db_user="litellm_rds", + ) + + +def test_entra_mint_calls_the_injected_provider_and_encodes_the_token(): + """A real compact JWT is already URL-safe, but the provider is an Azure SDK call + whose output we do not control, and an unencoded ``/`` or ``=`` in a password + silently truncates the connection URL.""" + auth = AzureEntraTokenAuth(token_provider=lambda: "head.pay/load+x=.sig") + + assert mint_database_token(auth, _endpoint()) == "head.pay%2Fload%2Bx%3D.sig" + + +def test_entra_mint_asks_the_provider_every_time(): + """A refresh must get a new token, not a cached one from construction time.""" + tokens = iter(["first", "second"]) + auth = AzureEntraTokenAuth(token_provider=lambda: next(tokens)) + + assert mint_database_token(auth, _endpoint()) == "first" + assert mint_database_token(auth, _endpoint()) == "second" + + +# --------------------------------------------------------------------------- +# Expiry parsing +# --------------------------------------------------------------------------- + + +def test_rds_expiry_reads_the_sigv4_query_params(): + token = "writer.aurora.local:5432/?Action=connect&X-Amz-Date=20260820T101500Z&X-Amz-Expires=900" + + assert parse_database_token_expiration(RdsIamTokenAuth(), token) == datetime(2026, 8, 20, 10, 30, 0) + + +@pytest.mark.parametrize( + "token", + [ + "no-query-params", + "host/?X-Amz-Date=20260820T101500Z", + "host/?X-Amz-Expires=900", + "host/?X-Amz-Date=not-a-date&X-Amz-Expires=900", + ], +) +def test_rds_expiry_returns_none_when_unreadable(token): + assert parse_database_token_expiration(RdsIamTokenAuth(), token) is None + + +@pytest.mark.parametrize("exp", [1787000000, 1787000001, 1787000012, 1787000123]) +def test_entra_expiry_decodes_the_jwt_exp_claim(exp): + """Parametrized over several ``exp`` values so the payload length lands on every + base64 padding remainder: the JWT payload is stripped of its ``=`` padding and has + to be re-padded before it can be decoded.""" + auth = AzureEntraTokenAuth(token_provider=lambda: "unused") + + parsed = parse_database_token_expiration(auth, _entra_token(exp)) + + assert parsed is not None + assert parsed.tzinfo is None + assert parsed == datetime.fromtimestamp(exp, tz=timezone.utc).replace(tzinfo=None) + + +@pytest.mark.parametrize( + "token", + [ + "not-a-jwt", + "only.two", + "head.{}.sig", + "head.bm90LWpzb24.sig", + f"head.{base64.urlsafe_b64encode(b'{}').decode()}.sig", + f"head.{base64.urlsafe_b64encode(json.dumps({'exp': 'soon'}).encode()).decode()}.sig", + ], +) +def test_entra_expiry_returns_none_when_unreadable(token): + """An unreadable expiry must degrade to the caller's fallback refresh interval + rather than blowing up the refresh loop.""" + auth = AzureEntraTokenAuth(token_provider=lambda: "unused") + + assert parse_database_token_expiration(auth, token) is None + + +# --------------------------------------------------------------------------- +# URL building and parsing +# --------------------------------------------------------------------------- + + +def test_build_url_encodes_a_upn_user_and_the_schema(): + endpoint = _endpoint(user="litellm@contoso.onmicrosoft.com", name="litellm db", schema="app/schema") + + assert endpoint.build_url("TOKEN") == ( + "postgresql://litellm%40contoso.onmicrosoft.com:TOKEN" + "@pg.postgres.database.azure.com:5432/litellm%20db?schema=app%2Fschema" + ) + + +@pytest.mark.parametrize( + ("field", "value"), + [ + ("user", "svc%40corp"), + ("name", "litellm%20db"), + ("schema", "app%2Fschema"), + ], +) +def test_build_url_leaves_an_already_encoded_component_alone(field, value): + """RDS IAM auth interpolated these raw, so pre-encoding was the only way to get an + ``@`` into ``DATABASE_USER``. Encoding again turns ``svc%40corp`` into + ``svc%2540corp``, which Postgres rejects with ``User `svc%40corp` was denied + access``, so an operator who did that on RDS breaks on upgrade.""" + url = _endpoint(**{field: value}).build_url("TOKEN") + + assert value in url + assert "%25" not in url + + +def test_build_url_inserts_the_token_verbatim(): + """Both providers hand the token back already in wire form, so re-encoding it here + would double-escape the password.""" + rds_token = "writer.aurora.local%3A5432%2F%3FAction%3Dconnect%26X-Amz-Date%3D20260820T101500Z" + + assert _endpoint().build_url(rds_token) == ( + f"postgresql://litellm:{rds_token}@pg.postgres.database.azure.com:5432/litellm_db" + ) + + +@pytest.mark.parametrize( + "endpoint", + [ + IAMEndpoint(host="h.example.com", port="5432", user="litellm", name="litellm_db"), + IAMEndpoint(host="h.example.com", port="6543", user="litellm@contoso.com", name="db", schema="public"), + IAMEndpoint(host="h.example.com", port="5432", user="u", name="litellm db", schema="app schema"), + ], +) +def test_build_url_and_parse_round_trip(endpoint): + assert parse_iam_endpoint_from_url(endpoint.build_url("TOKEN")) == endpoint + + +def test_parse_leaves_an_already_escaped_schema_alone(): + """``parse_qs`` unquotes query values itself, so unquoting again here would turn a + schema that legitimately contains ``%40`` into one containing ``@``.""" + url = "postgresql://u:TOKEN@h.example.com:5432/db?schema=raw%2540schema" + + assert parse_iam_endpoint_from_url(url).schema == "raw%40schema" + + +# --------------------------------------------------------------------------- +# Strategy resolution from the environment +# --------------------------------------------------------------------------- + + +def test_resolve_returns_none_when_neither_toggle_is_set(monkeypatch): + monkeypatch.delenv(IAM_TOKEN_DB_AUTH_ENV_VAR, raising=False) + monkeypatch.delenv(AZURE_POSTGRESQL_AUTH_ENV_VAR, raising=False) + + assert resolve_database_token_auth() is None + + +def test_resolve_returns_the_rds_strategy(monkeypatch): + monkeypatch.setenv(IAM_TOKEN_DB_AUTH_ENV_VAR, "true") + monkeypatch.delenv(AZURE_POSTGRESQL_AUTH_ENV_VAR, raising=False) + + assert resolve_database_token_auth() == RdsIamTokenAuth() + + +def test_resolve_returns_the_entra_strategy(monkeypatch): + monkeypatch.delenv(IAM_TOKEN_DB_AUTH_ENV_VAR, raising=False) + monkeypatch.setenv(AZURE_POSTGRESQL_AUTH_ENV_VAR, "true") + + with patch( + "litellm.secret_managers.get_azure_ad_token_provider.get_azure_ad_token_provider", + return_value=lambda: "ENTRA_TOKEN", + ): + auth = resolve_database_token_auth() + + assert isinstance(auth, AzureEntraTokenAuth) + assert auth.token_provider() == "ENTRA_TOKEN" + + +def test_resolve_raises_when_both_toggles_are_set(monkeypatch): + monkeypatch.setenv(IAM_TOKEN_DB_AUTH_ENV_VAR, "true") + monkeypatch.setenv(AZURE_POSTGRESQL_AUTH_ENV_VAR, "true") + + with pytest.raises(RuntimeError, match="can only come from one token source"): + resolve_database_token_auth() + + +def test_entra_provider_uses_the_ossrdbms_scope(): + """The wrong scope mints a token Azure Postgres rejects, so the scope is pinned.""" + with patch( + "litellm.secret_managers.get_azure_ad_token_provider.get_azure_ad_token_provider", + return_value=lambda: "ENTRA_TOKEN", + ) as get_provider: + build_azure_entra_token_provider() + + get_provider.assert_called_once_with(azure_scope="https://ossrdbms-aad.database.windows.net/.default") + assert AZURE_POSTGRESQL_SCOPE == "https://ossrdbms-aad.database.windows.net/.default" + + +# --------------------------------------------------------------------------- +# Toggle parsing +# --------------------------------------------------------------------------- + + +@pytest.mark.parametrize("value", ["true", "TRUE", " True ", "1", "yes", "y", "on", "t"]) +def test_every_truthy_spelling_enables_token_auth(monkeypatch, value): + """The settings model reads these toggles with pydantic (which accepts all of these) + while the refresh loop reads them here. When the two disagreed, `AZURE_POSTGRESQL_AUTH=1` + minted a token at startup and then never refreshed it, so the proxy died an hour in.""" + monkeypatch.setenv(AZURE_POSTGRESQL_AUTH_ENV_VAR, value) + monkeypatch.delenv(IAM_TOKEN_DB_AUTH_ENV_VAR, raising=False) + + with patch( + "litellm.secret_managers.get_azure_ad_token_provider.get_azure_ad_token_provider", + return_value=lambda: "ENTRA_TOKEN", + ): + assert isinstance(resolve_database_token_auth(), AzureEntraTokenAuth) + + +@pytest.mark.parametrize("value", ["", " ", "false", "False", "0", "no", "off", "F", "N"]) +def test_falsy_spellings_leave_token_auth_off(monkeypatch, value): + """An empty string is how a Kubernetes manifest spells 'off'.""" + monkeypatch.setenv(AZURE_POSTGRESQL_AUTH_ENV_VAR, value) + monkeypatch.setenv(IAM_TOKEN_DB_AUTH_ENV_VAR, value) + + assert resolve_database_token_auth() is None + + +@pytest.mark.parametrize("env_var", [IAM_TOKEN_DB_AUTH_ENV_VAR, AZURE_POSTGRESQL_AUTH_ENV_VAR]) +@pytest.mark.parametrize("value", ["enabled", "maybe", "TRUEE", "2"]) +def test_an_unreadable_toggle_is_a_startup_error(monkeypatch, env_var, value): + """Reading a typo as 'off' would quietly downgrade an operator who asked for token + auth to password auth, and the first sign of it is the server refusing the + connection. Pydantic rejected these before token auth had its own parser.""" + monkeypatch.setenv(env_var, value) + monkeypatch.delenv( + AZURE_POSTGRESQL_AUTH_ENV_VAR if env_var == IAM_TOKEN_DB_AUTH_ENV_VAR else IAM_TOKEN_DB_AUTH_ENV_VAR, + raising=False, + ) + + with pytest.raises(ValueError, match=env_var) as raised: + resolve_database_token_auth() + + assert value in str(raised.value) + + +def test_the_entra_provider_is_built_once_per_process(): + """Each build is another Azure credential with its own transport and token cache + that nothing closes, and the writer, the reader, and the refresh loop each ask.""" + with patch( + "litellm.secret_managers.get_azure_ad_token_provider.get_azure_ad_token_provider", + return_value=lambda: "ENTRA_TOKEN", + ) as get_provider: + assert build_azure_entra_token_provider() is build_azure_entra_token_provider() + + get_provider.assert_called_once() + + +def test_an_unparseable_rds_expiry_degrades_instead_of_raising(): + """This runs inside `PrismaWrapper.__getattr__`, so anything it raises turns every + database call into that error.""" + absurd = "https://host/?X-Amz-Date=20260820T101500Z&X-Amz-Expires=99999999999999999999" + + assert parse_database_token_expiration(RdsIamTokenAuth(), absurd) is None diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py index 8a551f749d0..8b9ecfbbeee 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/unified_guardrails/test_unified_guardrail.py @@ -1016,10 +1016,9 @@ class TestStreamingTransform: @pytest.mark.asyncio async def test_emit_streaming_http_error_a2a_yields_jsonrpc_chunk(self): - """The shared streaming error helper emits an in-stream JSON-RPC error for - A2A call types instead of raising.""" - import json - + """The shared streaming error helper emits an in-stream JSON-RPC error + object (not a pre-serialized string, which the A2A endpoint would frame as + a JSON string instead of an error object) for A2A call types.""" handler = UnifiedLLMGuardrails() exc = unified_module.HTTPException( status_code=400, @@ -1036,7 +1035,8 @@ class TestStreamingTransform: emitted.append(item) assert len(emitted) == 1 - payload = json.loads(emitted[0]) + payload = emitted[0] + assert isinstance(payload, dict) assert payload["error"]["message"] == "stream_transform_underflow" assert payload["id"] == "req-1" diff --git a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py index c4a0a62ef97..291ce732fc6 100644 --- a/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py +++ b/tests/test_litellm/proxy/guardrails/test_deferred_guardrail_logging.py @@ -625,7 +625,7 @@ class TestDeferredStreamingClosure: """If a guardrail raises HTTPException, the production _run_deferred_stream_guardrails must still fire logging and set guardrail_blocked in metadata.""" - from fastapi import HTTPException # noqa: local import for test isolation + from fastapi import HTTPException # local import for test isolation logging_called = False diff --git a/tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py b/tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py index c977223bfab..48f6b3ba2b9 100644 --- a/tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py +++ b/tests/test_litellm/proxy/guardrails/test_pillar_guardrails.py @@ -9,7 +9,7 @@ and following LiteLLM testing patterns and best practices. import importlib import os import sys -from typing import Dict +from typing import Any, Dict from unittest.mock import Mock, patch # Add parent directory to path for imports diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py index e333bf1e3fe..f29069f7f3c 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_v2_endpoints.py @@ -19,6 +19,7 @@ from litellm.proxy.management_endpoints.scim.scim_v2 import ( _apply_group_patch_updates, _extract_group_member_ids, _extract_ids_from_path_filter, + _handle_group_membership_changes, _handle_team_membership_changes, _parse_member_entries, _process_group_patch_operations, @@ -548,7 +549,12 @@ async def test_handle_existing_user_by_email_no_existing_user(mocker): @pytest.mark.asyncio async def test_handle_existing_user_by_email_existing_user_updated(mocker): - """Should rename the existing user, sync team roster, and return SCIMUser""" + """Should keep the existing user_id, sync team roster, and return SCIMUser + + Regression: a SCIM userName differing from the matched row's user_id used to + re-key the user row, orphaning virtual keys, team rosters, memberships and + spend logs that still referenced the old id. + """ existing_user = mocker.MagicMock() existing_user.user_id = "old-user-id" existing_user.user_email = "test@example.com" @@ -557,7 +563,7 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): existing_user.metadata = {"old": "data"} updated_user = { - "user_id": "new-user-id", + "user_id": "old-user-id", "user_email": "test@example.com", "user_alias": "New Name", "teams": ["new-team"], @@ -566,8 +572,8 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): mock_scim_user = SCIMUser( schemas=["urn:ietf:params:scim:schemas:core:2.0:User"], - id="new-user-id", - userName="new-user-id", + id="old-user-id", + userName="test@example.com", name=SCIMUserName(familyName="Name", givenName="New"), emails=[SCIMUserEmail(value="test@example.com")], ) @@ -605,13 +611,9 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): mock_prisma_client.db.litellm_usertable.find_first.assert_called_once_with(where={"user_email": "test@example.com"}) update_calls = mock_prisma_client.db.litellm_usertable.update.call_args_list - assert len(update_calls) == 2 + assert len(update_calls) == 1 assert update_calls[0].kwargs == { "where": {"user_id": "old-user-id"}, - "data": {"user_id": "new-user-id"}, - } - assert update_calls[1].kwargs == { - "where": {"user_id": "new-user-id"}, "data": { "user_email": "test@example.com", "user_alias": "New Name", @@ -621,7 +623,7 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): } mock_membership.assert_awaited_once_with( - user_id="new-user-id", + user_id="old-user-id", existing_teams=["old-team"], new_teams=["new-team"], raise_on_error=True, @@ -630,6 +632,57 @@ async def test_handle_existing_user_by_email_existing_user_updated(mocker): mock_transform.assert_called_once_with(updated_user) +@pytest.mark.asyncio +async def test_handle_existing_user_by_email_roster_changes_use_existing_user_id(mocker): + """Roster add/remove must be issued for the matched row's user_id, not the SCIM userName. + + Regression: the rename made removals run against the new id, so a roster still + holding the old id reported "User not found in team" and the stale entry survived. + """ + existing_user = mocker.MagicMock() + existing_user.user_id = "oidc-sub-123" + existing_user.user_email = "member@example.com" + existing_user.user_alias = "Member" + existing_user.teams = ["old-team"] + existing_user.metadata = {} + + mock_prisma_client = mocker.MagicMock() + mock_prisma_client.db = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable = mocker.MagicMock() + mock_prisma_client.db.litellm_usertable.find_first = AsyncMock(return_value=existing_user) + mock_prisma_client.db.litellm_usertable.update = AsyncMock(return_value={}) + + mock_team_member_add = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_add", + AsyncMock(), + ) + mock_team_member_delete = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_delete", + AsyncMock(), + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_user_to_scim_user", + AsyncMock(return_value=None), + ) + + new_user_request = NewUserRequest( + user_id="scim-username", + user_email="member@example.com", + user_alias="Member", + teams=["new-team"], + metadata={}, + auto_create_key=False, + ) + + await UserProvisionerHelpers.handle_existing_user_by_email( + prisma_client=mock_prisma_client, new_user_request=new_user_request + ) + + assert mock_team_member_add.await_args.kwargs["data"].member.user_id == "oidc-sub-123" + assert mock_team_member_delete.await_args.kwargs["data"].user_id == "oidc-sub-123" + assert mock_prisma_client.db.litellm_usertable.update.await_args.kwargs["where"] == {"user_id": "oidc-sub-123"} + + @pytest.mark.asyncio async def test_handle_existing_user_by_email_syncs_roster_and_dedups_teams(mocker): """Existing-email upsert must add the user to the team roster via the shared @@ -1245,6 +1298,10 @@ async def test_update_group_metadata_serialization_issue(mocker): "litellm.proxy.management_endpoints.scim.scim_v2.ScimTransformations.transform_litellm_team_to_scim_group", AsyncMock(return_value=mock_scim_group_response), ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.patch_team_membership", + AsyncMock(), + ) # Call the function that had the bug await update_group(group_id=group_id, group=scim_group) @@ -4401,3 +4458,89 @@ async def test_get_groups_members_are_typed_as_users(mocker): response = await get_groups(startIndex=1, count=10, filter=None) assert [m.type for m in response.Resources[0].members] == ["User"] + + +@pytest.mark.asyncio +async def test_resolve_group_member_ids_raises_when_creation_fails(mocker, scim_upsert_user_enabled): + """A member whose user row can neither be found nor created must fail the + request. Regression: the resolver silently dropped that member and the group + write reported success, so the IdP recorded the user as provisioned while the + team roster was missing them.""" + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=None), + ) + + with pytest.raises(HTTPException) as exc_info: + await _resolve_group_member_ids( + members=[SCIMMember(value="member-1")], + created_via="scim_group_membership", + prisma_client=_member_resolution_prisma(mocker, users=set(), teams=set()), + ) + + assert exc_info.value.status_code == 500 + assert "member-1" in str(exc_info.value.detail) + + +@pytest.mark.asyncio +async def test_resolve_group_member_ids_admits_member_created_concurrently(mocker, scim_upsert_user_enabled): + """When creation fails because a concurrent request already created the user, + the member is still admitted: the id resolves to a real user row, so failing + or dropping it would be wrong either way.""" + prisma_client = _member_resolution_prisma(mocker, users=set(), teams=set()) + prisma_client.db.litellm_usertable.find_unique = AsyncMock( + side_effect=[None, LiteLLM_UserTable(user_id="raced-user")] + ) + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2._create_user_if_not_exists", + AsyncMock(return_value=None), + ) + + result = await _resolve_group_member_ids( + members=[SCIMMember(value="raced-user")], + created_via="scim_group_membership", + prisma_client=prisma_client, + ) + + assert result.all_member_ids == ["raced-user"] + assert len(result.created_users) == 0 + + +@pytest.mark.asyncio +async def test_handle_group_membership_changes_propagates_add_failure(mocker): + """A genuine roster add failure must fail the group request so the IdP retries. + Regression: patch_team_membership ran with raise_on_error=False here, so a + failed team_member_add was logged and swallowed and the SCIM group sync + reported success with members missing from the team.""" + mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_add", + AsyncMock(side_effect=HTTPException(status_code=500, detail={"error": "db write failed"})), + ) + + with pytest.raises(HTTPException): + await _handle_group_membership_changes( + group_id="group-1", current_members=set(), final_members={"user-1"} + ) + + +@pytest.mark.asyncio +async def test_handle_group_membership_changes_already_in_team_is_noop(mocker): + """The strict path must keep treating an already-enrolled member as a no-op + and continue with the remaining members instead of failing the sync.""" + mock_team_member_add = mocker.patch( + "litellm.proxy.management_endpoints.scim.scim_v2.team_member_add", + AsyncMock( + side_effect=ProxyException( + message="already in team", + type=ProxyErrorTypes.team_member_already_in_team.value, + param=None, + code=400, + ) + ), + ) + + await _handle_group_membership_changes( + group_id="group-1", current_members=set(), final_members={"user-1", "user-2"} + ) + + assert mock_team_member_add.await_count == 2 diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index c901696e108..6165e869989 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -490,7 +490,7 @@ from litellm.proxy.management_endpoints.auto_router_endpoints import ( start_shadow_eval, stop_shadow_eval_job, ) -from litellm.types.management_endpoints.auto_router_endpoints import StartShadowEvalRequest +from litellm.types.management_endpoints.auto_router_endpoints import SHADOW_EVAL_TURN_VALVE, StartShadowEvalRequest VIEWER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, api_key="sk-view", user_id="viewer") NON_ADMIN = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-user", user_id="user") @@ -520,6 +520,7 @@ def _leg_record(**overrides: object) -> MagicMock: "judge_model": "anthropic/claude-sonnet-5", "shadow_percentage": 10.0, "max_turns": 200, + "max_budget": None, "created_at": datetime(2026, 8, 11, tzinfo=timezone.utc), "ends_at": datetime.now(timezone.utc) + timedelta(days=7), "stopped_at": None, @@ -549,11 +550,18 @@ def _shadow_prisma(legs=(), agg_rows=None, by_leg_rows=None, known_keys=("key-ha group read that matched on a leg id would come back empty.""" prisma = MagicMock() prisma.db.litellm_verificationtoken.find_many = AsyncMock(return_value=[_key_record(token) for token in known_keys]) + async def execute_raw(sql: str, *params: object): if "SET stopped_by" in sql: group = [row for row in stored if row.group_id == params[0]] counts = {row["job_id"]: row["attempt_count"] for row in prisma.attempt_rows} - sampling = any(row.stopped_at is None and counts.get(row.id, 0) < row.max_turns for row in group) + spends = {row["job_id"]: row["spend"] for row in prisma.attempt_rows} + sampling = any( + row.stopped_at is None + and counts.get(row.id, 0) < row.max_turns + and (row.max_budget is None or spends.get(row.id, 0.0) < row.max_budget) + for row in group + ) window_open = bool(group) and group[0].ends_at > datetime.now(timezone.utc) claimable = [row for row in group if row.stopped_by is None] if not (claimable and sampling and window_open): @@ -602,6 +610,7 @@ def _shadow_prisma(legs=(), agg_rows=None, by_leg_rows=None, known_keys=("key-ha "judge_model", "shadow_percentage", "max_turns", + "max_budget", "created_at", "ends_at", "stopped_at", @@ -639,7 +648,7 @@ def _start_request(**overrides: object) -> StartShadowEvalRequest: "shadow_percentage": 10.0, "judge_model": "anthropic/claude-sonnet-5", "duration_days": 7, - "max_turns": 200, + "max_budget": 5.0, } payload.update(overrides) return StartShadowEvalRequest.model_validate(payload) @@ -663,6 +672,9 @@ async def test_start_shadow_eval_writes_one_leg_per_key_in_one_statement(monkeyp assert "j.ends_at <= (NOW() AT TIME ZONE 'utc')" in sweep_sql assert "SET stopped_at = (NOW() AT TIME ZONE 'utc')" in sweep_sql assert ">= j.max_turns" in sweep_sql + assert "j.max_budget IS NOT NULL" in sweep_sql + assert ">= j.max_budget" in sweep_sql + assert "SUM(a.judge_cost + a.shadow_cost)" in sweep_sql assert "j.api_key_id = ANY($1::text[])" in sweep_sql assert sweep_keys == ["key-hash", "key-hash-2"] prisma.db.litellm_shadowevaljob.create_many.assert_awaited_once() @@ -670,15 +682,17 @@ async def test_start_shadow_eval_writes_one_leg_per_key_in_one_statement(monkeyp assert [row["api_key_id"] for row in rows] == ["key-hash", "key-hash-2"] assert len({frozenset((k, v) for k, v in row.items() if k != "api_key_id") for row in rows}) == 1 assert len({row["group_id"] for row in rows}) == 1 - assert all(row["max_turns"] == 200 and row["created_by"] == "admin" for row in rows) + assert all(row["max_turns"] == SHADOW_EVAL_TURN_VALVE and row["created_by"] == "admin" for row in rows) + assert all(row["max_budget"] == 5.0 for row in rows) assert all("status" not in row and "id" not in row for row in rows) assert response.job_id == rows[0]["group_id"] assert response.status == "running" assert response.judged_count is None - assert [(key.api_key_id, key.max_turns, key.key_alias) for key in response.keys] == [ - ("key-hash", 200, "prod-alpha"), - ("key-hash-2", 200, "prod-alpha"), + assert [(key.api_key_id, key.max_budget, key.key_alias) for key in response.keys] == [ + ("key-hash", 5.0, "prod-alpha"), + ("key-hash-2", 5.0, "prod-alpha"), ] + assert all(key.max_turns == SHADOW_EVAL_TURN_VALVE for key in response.keys) @pytest.mark.asyncio @@ -1041,10 +1055,10 @@ async def test_list_reads_completed_once_every_key_spends_its_budget(monkeypatch ] ) prisma.attempt_rows = [ - {"job_id": "leg-1", "attempt_count": 5}, - {"job_id": "leg-2", "attempt_count": 6}, - {"job_id": "leg-3", "attempt_count": 5}, - {"job_id": "leg-4", "attempt_count": 3}, + {"job_id": "leg-1", "attempt_count": 5, "spend": 0.0}, + {"job_id": "leg-2", "attempt_count": 6, "spend": 0.0}, + {"job_id": "leg-3", "attempt_count": 5, "spend": 0.0}, + {"job_id": "leg-4", "attempt_count": 3, "spend": 0.0}, ] monkeypatch.setattr(proxy_server, "prisma_client", prisma) @@ -1065,7 +1079,7 @@ async def test_recorded_operator_stop_outranks_budget_arithmetic(monkeypatch: py stamp = datetime.now(timezone.utc) prisma = _shadow_prisma(legs=[_leg_record(max_turns=5, stopped_at=stamp, stopped_by="admin")]) - prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 6}] + prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 6, "spend": 0.0}] monkeypatch.setattr(proxy_server, "prisma_client", prisma) jobs = await list_shadow_eval_jobs(VIEWER, api_key_id=None, limit=50) @@ -1085,7 +1099,7 @@ async def test_backfilled_legacy_stop_never_reads_as_completion(monkeypatch: pyt prisma = _shadow_prisma( legs=[_leg_record(max_turns=5, stopped_at=datetime.now(timezone.utc), stopped_by="unknown")] ) - prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 6}] + prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 6, "spend": 0.0}] monkeypatch.setattr(proxy_server, "prisma_client", prisma) jobs = await list_shadow_eval_jobs(VIEWER, api_key_id=None, limit=50) @@ -1108,12 +1122,39 @@ def test_stopped_by_migration_backfills_every_job_that_displayed_stopped(): assert "WHERE stopped_at IS NOT NULL AND ends_at > (NOW() AT TIME ZONE 'utc')" in sql +def test_a_start_request_still_sending_max_turns_is_rejected_not_silently_defaulted(): + """Pydantic ignores unknown fields, so without the explicit rejection a caller still + sending the retired turn budget would silently run on the default dollar budget.""" + with pytest.raises(ValidationError, match="max_budget"): + _start_request(max_turns=200) + + +def test_max_budget_migration_is_additive_and_leaves_legacy_rows_null(): + """max_budget stays NULL on pre-migration rows so they keep the turn budget they were + configured with, and shadow_cost defaults to 0 so old rows price as judge-only.""" + import litellm_proxy_extras + + sql = ( + Path(litellm_proxy_extras.__file__).parent + / "migrations" + / "20260819000000_shadow_eval_max_budget" + / "migration.sql" + ).read_text() + assert 'ALTER TABLE "LiteLLM_ShadowEvalJob" ADD COLUMN "max_budget" DOUBLE PRECISION' in sql + assert ( + 'ALTER TABLE "LiteLLM_ShadowEvalAttempt" ADD COLUMN "shadow_cost" DOUBLE PRECISION NOT NULL DEFAULT 0' + in sql + ) + assert "UPDATE" not in sql + assert "DROP" not in sql + + @pytest.mark.asyncio async def test_stop_rejects_a_job_that_already_spent_its_budget(monkeypatch: pytest.MonkeyPatch): import litellm.proxy.proxy_server as proxy_server prisma = _shadow_prisma(legs=[_leg_record(max_turns=3)]) - prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 3}] + prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 3, "spend": 0.0}] monkeypatch.setattr(proxy_server, "prisma_client", prisma) with pytest.raises(HTTPException) as exhausted: @@ -1123,6 +1164,71 @@ async def test_stop_rejects_a_job_that_already_spent_its_budget(monkeypatch: pyt prisma.db.litellm_shadowevaljob.update_many.assert_not_called() +@pytest.mark.asyncio +async def test_list_reads_completed_once_every_key_spends_its_dollar_budget(monkeypatch: pytest.MonkeyPatch): + """A spend-budgeted job completes on dollars, not turns: every key's recorded shadow + plus judge spend reaching max_budget reads completed long before the turn valve, while + one key with budget left keeps the whole job running.""" + import litellm.proxy.proxy_server as proxy_server + + prisma = _shadow_prisma( + legs=[ + _leg_record(max_turns=SHADOW_EVAL_TURN_VALVE, max_budget=1.0), + _leg_record(id="leg-2", api_key_id="key-hash-2", max_turns=SHADOW_EVAL_TURN_VALVE, max_budget=1.0), + _leg_record( + id="leg-3", group_id="job-2", api_key_id="key-hash", max_turns=SHADOW_EVAL_TURN_VALVE, max_budget=1.0 + ), + ] + ) + prisma.attempt_rows = [ + {"job_id": "leg-1", "attempt_count": 40, "spend": 1.0}, + {"job_id": "leg-2", "attempt_count": 55, "spend": 1.25}, + {"job_id": "leg-3", "attempt_count": 40, "spend": 0.99}, + ] + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + + jobs = await list_shadow_eval_jobs(VIEWER, api_key_id=None, limit=50) + + by_id = {job.job_id: job for job in jobs} + assert by_id["job-1"].status == "completed" + assert by_id["job-2"].status == "running" + assert {key.api_key_id: key.spend for key in by_id["job-1"].keys} == {"key-hash": 1.0, "key-hash-2": 1.25} + assert all(key.max_budget == 1.0 for key in by_id["job-1"].keys) + + +@pytest.mark.asyncio +async def test_stop_rejects_a_job_whose_dollar_budget_is_spent(monkeypatch: pytest.MonkeyPatch): + import litellm.proxy.proxy_server as proxy_server + + prisma = _shadow_prisma(legs=[_leg_record(max_turns=SHADOW_EVAL_TURN_VALVE, max_budget=0.5)]) + prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 7, "spend": 0.5}] + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + + with pytest.raises(HTTPException) as exhausted: + await stop_shadow_eval_job("job-1", ADMIN) + assert exhausted.value.status_code == 400 + assert "completed" in exhausted.value.detail + prisma.db.litellm_shadowevaljob.update_many.assert_not_called() + + +@pytest.mark.asyncio +async def test_legacy_jobs_without_a_dollar_budget_stay_turn_gated(monkeypatch: pytest.MonkeyPatch): + """A job from before spend budgets existed carries max_budget NULL: recorded spend + can never complete it, only its own max_turns can, so migration changes nothing about + what it was configured to do.""" + import litellm.proxy.proxy_server as proxy_server + + prisma = _shadow_prisma(legs=[_leg_record(max_turns=200, max_budget=None)]) + prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 40, "spend": 250.0}] + monkeypatch.setattr(proxy_server, "prisma_client", prisma) + + jobs = await list_shadow_eval_jobs(VIEWER, api_key_id=None, limit=50) + + assert jobs[0].status == "running" + assert jobs[0].keys[0].max_budget is None + assert jobs[0].keys[0].spend == 250.0 + + @pytest.mark.asyncio async def test_shadow_eval_responses_name_every_shadowed_key(monkeypatch: pytest.MonkeyPatch): import litellm.proxy.proxy_server as proxy_server @@ -1167,6 +1273,9 @@ async def test_stop_shadow_eval_stops_every_unstopped_leg_and_rejects_non_runnin assert "WHERE group_id = $1 AND stopped_by IS NULL" in stop_sql assert "ends_at > (NOW() AT TIME ZONE 'utc')" in stop_sql assert ") < k.max_turns" in stop_sql + assert "k.max_budget IS NULL" in stop_sql + assert ") < k.max_budget" in stop_sql + assert "SUM(a.judge_cost + a.shadow_cost)" in stop_sql assert (stop_group, stop_operator) == ("job-1", "admin") assert datetime.fromisoformat(stop_stamp).tzinfo is None assert prisma.db.execute_raw.await_count == 1 @@ -1357,7 +1466,7 @@ async def test_a_stop_racing_the_last_budgeted_attempt_reports_completed_not_sto import litellm.proxy.proxy_server as proxy_server prisma = _shadow_prisma(legs=[_leg_record(max_turns=2)]) - prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 2}] + prisma.attempt_rows = [{"job_id": "leg-1", "attempt_count": 2, "spend": 0.0}] monkeypatch.setattr(proxy_server, "prisma_client", prisma) with pytest.raises(HTTPException) as exc: diff --git a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py index 7dfd99dfa53..da8fc760787 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_common_utils.py +++ b/tests/test_litellm/proxy/management_endpoints/test_common_utils.py @@ -30,6 +30,7 @@ from litellm.proxy.management_endpoints.common_utils import ( _user_has_admin_view, admin_can_invite_user, ) +from litellm.proxy.management_endpoints.common_utils import _has_non_empty_value class TestUpdateMetadataFieldsEmptyCollections: @@ -977,3 +978,145 @@ class TestUpdateMetadataFieldMove: _update_metadata_fields(updated_kv) assert "guardrails" not in updated_kv assert updated_kv["metadata"]["guardrails"] == ["g1"] + + +class TestHasNonEmptyValue: + """Tests for the _has_non_empty_value helper.""" + + def test_none_is_empty(self): + assert _has_non_empty_value(None) is False + + def test_empty_list_is_empty(self): + assert _has_non_empty_value([]) is False + + def test_empty_string_is_empty(self): + assert _has_non_empty_value("") is False + + def test_blank_string_is_empty(self): + assert _has_non_empty_value(" ") is False + + def test_non_empty_list_has_value(self): + assert _has_non_empty_value(["policy-a"]) is True + + def test_non_empty_string_has_value(self): + assert _has_non_empty_value("30d") is True + + def test_dict_has_value(self): + assert _has_non_empty_value({"key": "val"}) is True + + def test_empty_dict_has_value(self): + # empty dict is not None/list/str, so it counts as non-empty + assert _has_non_empty_value({}) is True + + +class TestUpdateMetadataFieldsPremiumCheck: + """ + Tests that _update_metadata_fields skips premium user checks for empty + values but still enforces them for real values. + + Issue: The UI sends the full form on every team update, including premium + fields like `policies: []`. The backend was treating these empty values + as premium feature usage and returning 403. + """ + + @patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check", + side_effect=Exception("Should not be called"), + ) + def test_empty_policies_skips_premium_check(self, mock_check): + """policies: [] should NOT trigger premium user check.""" + updated_kv = { + "team_id": "team-123", + "team_alias": "my-team", + "policies": [], + } + _update_metadata_fields(updated_kv) + mock_check.assert_not_called() + + @patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check", + side_effect=Exception("Should not be called"), + ) + def test_empty_guardrails_skips_premium_check(self, mock_check): + """guardrails: [] should NOT trigger premium user check.""" + updated_kv = { + "team_id": "team-123", + "guardrails": [], + } + _update_metadata_fields(updated_kv) + mock_check.assert_not_called() + + @patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check", + side_effect=Exception("Should not be called"), + ) + def test_empty_string_team_member_key_duration_skips_premium_check( + self, mock_check + ): + """team_member_key_duration: '' should NOT trigger premium user check.""" + updated_kv = { + "team_id": "team-123", + "team_member_key_duration": "", + } + _update_metadata_fields(updated_kv) + mock_check.assert_not_called() + + @patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check", + side_effect=Exception("Should not be called"), + ) + def test_full_ui_payload_with_empty_premium_fields_skips_premium_check( + self, mock_check + ): + """A realistic UI payload with all empty premium fields should not 403.""" + updated_kv = { + "team_id": "team-123", + "team_alias": "renamed-team", + "models": ["gpt-4o"], + "max_budget": 200, + "policies": [], + "guardrails": [], + "logging": [], + "team_member_key_duration": "", + "prompts": [], + } + _update_metadata_fields(updated_kv) + mock_check.assert_not_called() + + @patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check", + ) + def test_non_empty_policies_triggers_premium_check(self, mock_check): + """policies: ['real-policy'] SHOULD trigger premium user check.""" + updated_kv = { + "team_id": "team-123", + "policies": ["real-policy"], + } + _update_metadata_fields(updated_kv) + mock_check.assert_called() + + @patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check", + ) + def test_non_empty_guardrails_triggers_premium_check(self, mock_check): + """guardrails: ['my-guardrail'] SHOULD trigger premium user check.""" + updated_kv = { + "team_id": "team-123", + "guardrails": ["my-guardrail"], + } + _update_metadata_fields(updated_kv) + mock_check.assert_called() + + @patch( + "litellm.proxy.management_endpoints.common_utils._premium_user_check", + ) + def test_non_empty_team_member_key_duration_triggers_premium_check( + self, mock_check + ): + """team_member_key_duration: '30d' SHOULD trigger premium user check.""" + updated_kv = { + "team_id": "team-123", + "team_member_key_duration": "30d", + } + _update_metadata_fields(updated_kv) + mock_check.assert_called() diff --git a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py index 2d54a391cf0..0767288d0bc 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_team_endpoints.py @@ -8454,9 +8454,7 @@ async def test_get_team_daily_activity_member_with_permission_sees_all_spend( hasattr(mock_db_client.db.litellm_verificationtoken, "find_many") and mock_db_client.db.litellm_verificationtoken.find_many.called ): - assert ( - False - ), "API keys should not be fetched for members with /team/daily/activity permission" + pytest.fail("API keys should not be fetched for members with /team/daily/activity permission") @pytest.mark.asyncio @@ -8808,7 +8806,7 @@ async def test_get_team_daily_activity_team_admin_sees_all_spend(mock_db_client) and mock_db_client.db.litellm_verificationtoken.find_many.called ): # If it was called, that's unexpected for admin users - assert False, "API keys should not be fetched for team admin users" + pytest.fail("API keys should not be fetched for team admin users") @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py index 80b7e25e07d..66cb07ef2c0 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ui_sso.py @@ -1606,6 +1606,294 @@ async def test_get_generic_sso_response_with_empty_headers(): assert result == mock_sso_response +@pytest.mark.asyncio +async def test_get_generic_sso_response_includes_token_claims_when_enabled(monkeypatch): + import jwt as pyjwt + + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + from litellm.proxy._types import LitellmUserRoles + + mock_request = MagicMock(spec=Request) + mock_jwt_handler = MagicMock(spec=JWTHandler) + mock_sso_jwt_handler = MagicMock(spec=JWTHandler) + mock_sso_jwt_handler.get_all_jwt_team_ids.return_value = ["team-from-userinfo"] + mock_sso_jwt_handler.get_team_ids_from_jwt.return_value = [] + + userinfo = { + "sub": "subject-only", + "groups": ["admins"], + "access_token": "", + } + access_token = pyjwt.encode( + { + "upn": "token-user@example.com", + "email": "token-user@example.com", + "given_name": "Token", + "family_name": "User", + "display_name": "Token User", + }, + "test-secret", + algorithm="HS256", + ) + mock_sso_instance = MagicMock() + mock_sso_instance.access_token = access_token + mock_sso_instance.id_token = None + + def fake_create_provider(*, response_convertor, **_kwargs): + mock_sso_instance.verify_and_process = AsyncMock( + side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object()) + ) + return MagicMock(return_value=mock_sso_instance) + + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo") + monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true") + monkeypatch.setenv("GENERIC_USER_ID_ATTRIBUTE", "upn") + monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email") + monkeypatch.setenv("GENERIC_USER_FIRST_NAME_ATTRIBUTE", "given_name") + monkeypatch.setenv("GENERIC_USER_LAST_NAME_ATTRIBUTE", "family_name") + monkeypatch.setenv("GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "display_name") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['admins']}") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "groups") + + with patch("fastapi_sso.sso.base.DiscoveryDocument"): + with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider): + result, received_response, _, _ = await get_generic_sso_response( + request=mock_request, + jwt_handler=mock_jwt_handler, + generic_client_id="test-client", + redirect_url="http://test.com/callback", + sso_jwt_handler=mock_sso_jwt_handler, + ) + + assert isinstance(result, CustomOpenID) + assert result.id == "token-user@example.com" + assert result.email == "token-user@example.com" + assert result.first_name == "Token" + assert result.last_name == "User" + assert result.display_name == "Token User" + assert result.team_ids == ["team-from-userinfo"] + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + assert received_response is not None + assert "access_token" not in received_response + assert "id_token" not in received_response + assert "refresh_token" not in received_response + + +@pytest.mark.asyncio +async def test_get_generic_sso_response_does_not_include_token_claims_when_disabled(monkeypatch): + import jwt as pyjwt + + from litellm.proxy._types import LitellmUserRoles + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + + mock_request = MagicMock(spec=Request) + mock_jwt_handler = MagicMock(spec=JWTHandler) + mock_sso_jwt_handler = MagicMock(spec=JWTHandler) + mock_sso_jwt_handler.get_all_jwt_team_ids.return_value = ["team-from-userinfo"] + mock_sso_jwt_handler.get_team_ids_from_jwt.return_value = [] + access_token = pyjwt.encode({"upn": "token-user@example.com"}, "test-secret", algorithm="HS256") + userinfo = {"sub": "subject-only", "groups": ["admins"], "access_token": ""} + mock_sso_instance = MagicMock() + mock_sso_instance.access_token = access_token + mock_sso_instance.id_token = None + + def fake_create_provider(*, response_convertor, **_kwargs): + mock_sso_instance.verify_and_process = AsyncMock( + side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object()) + ) + return MagicMock(return_value=mock_sso_instance) + + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo") + monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "false") + monkeypatch.setenv("GENERIC_USER_ID_ATTRIBUTE", "upn") + monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email") + monkeypatch.setenv("GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "display_name") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_ROLES", "{'proxy_admin': ['admins']}") + monkeypatch.setenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", "groups") + + with patch("fastapi_sso.sso.base.DiscoveryDocument"): + with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider): + result, received_response, _, _ = await get_generic_sso_response( + request=mock_request, + jwt_handler=mock_jwt_handler, + generic_client_id="test-client", + redirect_url="http://test.com/callback", + sso_jwt_handler=mock_sso_jwt_handler, + ) + + assert isinstance(result, CustomOpenID) + assert result.id is None + assert result.email is None + assert result.display_name is None + assert result.team_ids == ["team-from-userinfo"] + assert result.user_role == LitellmUserRoles.PROXY_ADMIN + assert received_response == {"sub": "subject-only", "groups": ["admins"]} + + +def test_merge_sso_token_claims_precedence_and_invalid_tokens(): + import jwt as pyjwt + + from litellm.proxy.management_endpoints.ui_sso import _merge_sso_token_claims + + id_token = pyjwt.encode( + {"preferred_username": "id-user", "email": "id@example.com", "id_only": "id-value"}, + "test-secret", + algorithm="HS256", + ) + access_token = pyjwt.encode( + {"preferred_username": "access-user", "email": "access@example.com", "access_only": "access-value"}, + "test-secret", + algorithm="HS256", + ) + + merged = _merge_sso_token_claims( + userinfo={"preferred_username": "userinfo-user", "email": None, "userinfo_only": "userinfo-value"}, + id_token=id_token, + access_token=access_token, + ) + + assert merged["preferred_username"] == "userinfo-user" + assert merged["email"] == "id@example.com" + assert merged["id_only"] == "id-value" + assert merged["access_only"] == "access-value" + + userinfo_only = _merge_sso_token_claims( + userinfo={"sub": "userinfo-user", "email": "userinfo@example.com"}, + id_token=pyjwt.encode({}, "test-secret", algorithm="HS256"), + access_token="opaque-access-token", + ) + + assert userinfo_only == {"sub": "userinfo-user", "email": "userinfo@example.com"} + + +@pytest.mark.asyncio +async def test_get_generic_sso_response_pkce_merges_token_claims_and_excludes_credentials(monkeypatch): + """The real PKCE path merges access-token claims and keeps bearer credentials out of received_response. + + Only the PKCE verifier cache and the HTTP transport are injected, so + prepare_token_exchange_parameters, _pkce_token_exchange and the claim merge all run for real. + """ + import jwt as pyjwt + from starlette.requests import Request as StarletteRequest + + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + + access_token = pyjwt.encode( + {"sub": "token-user", "email": "token-user@example.com"}, "test-secret", algorithm="HS256" + ) + request = StarletteRequest( + { + "type": "http", + "method": "GET", + "path": "/sso/callback", + "query_string": b"code=test-code&state=test-state", + "headers": [(b"cookie", b"litellm_oauth_state=test-state")], + } + ) + + pkce_cache = MagicMock(redis_cache=None) + pkce_cache.async_get_cache = AsyncMock(return_value={"code_verifier": "test-code-verifier"}) + pkce_cache.async_delete_cache = AsyncMock() + + token_endpoint_response = MagicMock(status_code=200) + token_endpoint_response.json.return_value = { + "access_token": access_token, + "id_token": "id-token-secret", + "refresh_token": "refresh-token-secret", + } + token_client = MagicMock() + token_client.post = AsyncMock(return_value=token_endpoint_response) + + userinfo_endpoint_response = MagicMock(status_code=200) + userinfo_endpoint_response.json.return_value = {"sub": "userinfo-user"} + userinfo_client = MagicMock() + userinfo_client.get = AsyncMock(return_value=userinfo_endpoint_response) + + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo") + monkeypatch.setenv("GENERIC_CLIENT_USE_PKCE", "true") + monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true") + monkeypatch.setenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email") + + with ( + patch("litellm.proxy.proxy_server.redis_usage_cache", None), + patch("litellm.proxy.proxy_server.user_api_key_cache", pkce_cache), + patch( + "litellm.proxy.management_endpoints.ui_sso.get_async_httpx_client", + side_effect=[token_client, userinfo_client], + ), + ): + result, received_response, _, _ = await get_generic_sso_response( + request=request, + jwt_handler=MagicMock(spec=JWTHandler), + generic_client_id="test-client", + redirect_url="http://test.com/callback", + sso_jwt_handler=None, + ) + + # The real token exchange ran: it forwarded the cached verifier to the token endpoint. + assert token_client.post.await_args.kwargs["data"]["code_verifier"] == "test-code-verifier" + # UserInfo wins for sub; email exists only on the access token, so the merge must supply it. + assert isinstance(result, CustomOpenID) + assert result.email == "token-user@example.com" + assert received_response == {"sub": "userinfo-user", "email": "token-user@example.com"} + pkce_cache.async_delete_cache.assert_awaited_once_with(key="pkce_verifier:test-state") + + +@pytest.mark.asyncio +async def test_get_generic_sso_response_ignores_opaque_and_empty_token_claims(monkeypatch): + import jwt as pyjwt + + from litellm.proxy.management_endpoints.ui_sso import get_generic_sso_response + + mock_request = MagicMock(spec=Request) + mock_jwt_handler = MagicMock(spec=JWTHandler) + userinfo = { + "preferred_username": "userinfo-user", + "email": "userinfo@example.com", + "sub": "User Info", + } + mock_sso_instance = MagicMock() + mock_sso_instance.access_token = "opaque-access-token" + mock_sso_instance.id_token = pyjwt.encode({}, "test-secret", algorithm="HS256") + + def fake_create_provider(*, response_convertor, **_kwargs): + mock_sso_instance.verify_and_process = AsyncMock( + side_effect=lambda *_args, **_kwargs: response_convertor(userinfo, object()) + ) + return MagicMock(return_value=mock_sso_instance) + + monkeypatch.setenv("GENERIC_CLIENT_SECRET", "test-secret") + monkeypatch.setenv("GENERIC_AUTHORIZATION_ENDPOINT", "https://auth.example.com/auth") + monkeypatch.setenv("GENERIC_TOKEN_ENDPOINT", "https://auth.example.com/token") + monkeypatch.setenv("GENERIC_USERINFO_ENDPOINT", "https://auth.example.com/userinfo") + monkeypatch.setenv("GENERIC_INCLUDE_TOKEN_CLAIMS", "true") + + with patch("fastapi_sso.sso.base.DiscoveryDocument"): + with patch("fastapi_sso.sso.generic.create_provider", side_effect=fake_create_provider): + result, received_response, _, _ = await get_generic_sso_response( + request=mock_request, + jwt_handler=mock_jwt_handler, + generic_client_id="test-client", + redirect_url="http://test.com/callback", + sso_jwt_handler=None, + ) + + assert isinstance(result, CustomOpenID) + assert result.id == "userinfo-user" + assert result.email == "userinfo@example.com" + assert result.display_name == "User Info" + assert received_response == userinfo + + class TestCLISSOCallbackFunction: """Test the cli_sso_callback function specifically""" @@ -1664,7 +1952,7 @@ class TestAuthCallbackRouting: key_id = cli_state.split(":", 1)[1] assert key_id == "cli-test1234567890" else: - assert False, "CLI state should have been detected" + pytest.fail("CLI state should have been detected") def test_non_cli_state_routing(self): """Test that non-CLI states don't trigger CLI routing""" diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 58465772b3b..d8047cca728 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -2617,6 +2617,33 @@ async def test_ProxyConfig__reschedule_spend_log_cleanup_job_invalid_cron(monkey assert fake_scheduler.add_job.call_count == 0 +@pytest.mark.asyncio +async def test_ProxyConfig__reschedule_spend_log_cleanup_job_health_check_retention(monkeypatch): + fake_scheduler = MagicMock() + monkeypatch.setattr("litellm.proxy.proxy_server.scheduler", fake_scheduler) + monkeypatch.setattr( + "litellm.proxy.proxy_server.general_settings", + {"maximum_health_check_retention_period": "30d"}, + ) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + pc = ProxyConfig() + await pc._reschedule_spend_log_cleanup_job() + assert fake_scheduler.add_job.call_count == 1 + assert fake_scheduler.add_job.call_args.kwargs["id"] == "spend_log_cleanup_job" + + +@pytest.mark.asyncio +async def test_ProxyConfig__update_general_settings_updates_health_check_retention(monkeypatch): + settings = {} + monkeypatch.setattr("litellm.proxy.proxy_server.general_settings", settings) + pc = ProxyConfig() + reschedule = AsyncMock() + monkeypatch.setattr(pc, "_reschedule_spend_log_cleanup_job", reschedule) + await pc._update_general_settings({"maximum_health_check_retention_period": "30d"}) + assert settings["maximum_health_check_retention_period"] == "30d" + reschedule.assert_awaited_once() + + # --------------------------------------------------------------------------- # ProxyConfig._update_general_settings # --------------------------------------------------------------------------- diff --git a/tests/test_litellm/proxy/proxy_server/test_routes_utils.py b/tests/test_litellm/proxy/proxy_server/test_routes_utils.py index c6070437d35..f39192b171b 100644 --- a/tests/test_litellm/proxy/proxy_server/test_routes_utils.py +++ b/tests/test_litellm/proxy/proxy_server/test_routes_utils.py @@ -8,6 +8,7 @@ Pins (PR2): from __future__ import annotations +import asyncio from unittest.mock import AsyncMock, MagicMock import pytest @@ -56,6 +57,31 @@ def test_token_counter_happy_path(client, auth_as, patched_token_counter): } +def test_token_counter_counts_off_the_event_loop(client, auth_as, patched_token_counter, monkeypatch): + """ + A large prompt must not stall the proxy: the count runs in a worker thread, where there + is no running event loop, rather than on the loop serving other requests. + """ + counted_off_loop = [] + + def recording_counter(**kwargs): + try: + asyncio.get_running_loop() + counted_off_loop.append(False) + except RuntimeError: + counted_off_loop.append(True) + return 7 + + monkeypatch.setattr(litellm, "token_counter", recording_counter) + + with auth_as(): + response = client.post("/utils/token_counter", json={"model": "gpt-4", "prompt": "Hi there"}) + + assert response.status_code == 200 + assert response.json()["total_tokens"] == 7 + assert counted_off_loop == [True] + + def test_token_counter_missing_input_returns_400( client, auth_as, patched_token_counter ): diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 7052e050806..15a3e6609f0 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -2659,7 +2659,7 @@ class TestSpendLogsPayload: payload, expected_payload, ignore_keys=ignored_keys ) if differences: - assert False, f"Dictionary mismatch: {differences}" + pytest.fail(f"Dictionary mismatch: {differences}") def mock_anthropic_response(*args, **kwargs): mock_response = MagicMock() @@ -2755,7 +2755,7 @@ class TestSpendLogsPayload: payload, expected_payload, ignore_keys=ignored_keys ) if differences: - assert False, f"Dictionary mismatch: {differences}" + pytest.fail(f"Dictionary mismatch: {differences}") @pytest.mark.asyncio async def test_spend_logs_payload_success_log_with_router(self, monkeypatch): @@ -2849,7 +2849,7 @@ class TestSpendLogsPayload: payload, expected_payload, ignore_keys=ignored_keys ) if differences: - assert False, f"Dictionary mismatch: {differences}" + pytest.fail(f"Dictionary mismatch: {differences}") def _compare_nested_dicts( diff --git a/tests/test_litellm/proxy/test_budget_reservation.py b/tests/test_litellm/proxy/test_budget_reservation.py index 34adb4d2091..133b53bb18d 100644 --- a/tests/test_litellm/proxy/test_budget_reservation.py +++ b/tests/test_litellm/proxy/test_budget_reservation.py @@ -1,4 +1,6 @@ import asyncio +import threading +from collections.abc import Mapping from datetime import datetime, timedelta, timezone from unittest.mock import AsyncMock, MagicMock, patch @@ -19,6 +21,8 @@ from litellm.proxy._types import ( ) from litellm.proxy.common_request_processing import ProxyBaseLLMRequestProcessing from litellm.proxy.spend_tracking.budget_reservation import ( + TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS, + _approximate_input_size, estimate_request_max_cost, get_budget_window_start, invalidate_budget_reservation_counters, @@ -2583,3 +2587,255 @@ async def test_streaming_slow_path_processes_and_yields_chunk(spend_counter_stat assert received == [{"content": "hi"}] streaming_logging_obj.async_post_call_streaming_hook.assert_awaited_once() + + +def _tiered_router() -> Router: + return Router( + model_list=[ + { + "model_name": "dashscope/qwen3-max", + "litellm_params": {"model": "dashscope/qwen3-max", "api_key": "sk-fake"}, + "model_info": { + "max_input_tokens": 258048, + "max_output_tokens": 65536, + "tiered_pricing": [ + { + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 6e-06, + "range": [0, 32000], + }, + { + "input_cost_per_token": 2.4e-06, + "output_cost_per_token": 1.2e-05, + "range": [32000, 128000], + }, + ], + }, + } + ] + ) + + +def _body_with_content_size(model: str, content_chars: int) -> dict: + return { + "model": model, + "messages": [{"role": "user", "content": "token " * (content_chars // 6)}], + "max_tokens": 10, + } + + +@pytest.mark.asyncio +async def test_reservation_tokenizes_the_prompt_once(spend_counter_state): + """Tokenizing is the reservation path's dominant CPU cost, so a request is + tokenized once no matter how many cost estimates and pricing candidates it + is priced against. The max-cost and input-cost estimates each used to + re-tokenize the prompt, once per tiered-pricing candidate.""" + _, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth( + token="key-tokenize-once", spend=0.0, max_budget=100.0 + ) + request_body = _body_with_content_size("dashscope/qwen3-max", 600) + real_token_counter = litellm.token_counter + calls = [] + + def counting_token_counter(**kwargs): + calls.append(kwargs) + return real_token_counter(**kwargs) + + with patch.object(litellm, "token_counter", counting_token_counter): + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=_tiered_router(), + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert reservation["reserved_cost"] > 0 + assert reservation["input_cost"] > 0 + assert len(calls) == 1 + + +@pytest.mark.asyncio +async def test_large_prompt_is_tokenized_off_the_event_loop(spend_counter_state): + """Counting a large prompt inline blocks the event loop for the whole count, + stalling every other request the worker is serving.""" + _, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth(token="key-offloaded", spend=0.0, max_budget=100.0) + request_body = _body_with_content_size( + "gpt-4o-mini", TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS + 6000 + ) + threads = [] + + def recording_token_counter(**kwargs): + threads.append(threading.current_thread()) + return 1000 + + with patch.object(litellm, "token_counter", recording_token_counter): + reservation = await reserve_budget_for_request( + request_body=request_body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert threads + assert all(thread is not threading.main_thread() for thread in threads) + + +def _values_only_size(value: object) -> int: + """The keys-ignoring walk the fixture below is sized to defeat""" + if isinstance(value, Mapping): + return sum(_values_only_size(item) for item in value.values()) + if isinstance(value, (list, tuple)): + return sum(_values_only_size(item) for item in value) + return len(value) if isinstance(value, str) else 0 + + +_TOOL_PROPERTY_NAME_PREFIX = "service_metric_name_segment_" * 3 + + +def _key_heavy_tool(index: int) -> dict: + return { + "type": "function", + "function": { + "name": f"lookup_service_metric_{index}", + "parameters": { + "type": "object", + "properties": { + f"{_TOOL_PROPERTY_NAME_PREFIX}{index}_{field}": {"type": "string"} + for field in range(24) + }, + }, + }, + } + + +def _body_with_key_heavy_tool_schema(model: str) -> dict: + """A tool schema whose bulk is property names rather than property values""" + return { + "model": model, + "messages": [{"role": "user", "content": "which service is slow?"}], + "tools": [_key_heavy_tool(index) for index in range(24)], + "max_tokens": 10, + } + + +@pytest.mark.asyncio +async def test_large_tool_schema_is_tokenized_off_the_event_loop(spend_counter_state): + """Tool-schema property names are tokenized like any other text. Sizing a + request by its values alone hides a large schema below the threshold, so it + gets counted inline and stalls the loop the threshold exists to spare.""" + body = _body_with_key_heavy_tool_schema("gpt-4o-mini") + assert _values_only_size(body["tools"]) < TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS + assert _approximate_input_size(body) >= TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS + + _, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth(token="key-tool-schema", spend=0.0, max_budget=100.0) + threads = [] + + def recording_token_counter(**kwargs): + threads.append(threading.current_thread()) + return 1000 + + with patch.object(litellm, "token_counter", recording_token_counter): + reservation = await reserve_budget_for_request( + request_body=body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert threads + assert all(thread is not threading.main_thread() for thread in threads) + + +@pytest.mark.asyncio +async def test_large_tool_choice_is_tokenized_off_the_event_loop(spend_counter_state): + """tool_choice is handed to the tokenizer alongside the messages, so a + request is only sized correctly if the heuristic covers it too.""" + body = { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": "which service is slow?"}], + "tool_choice": { + "type": "function", + "function": {"name": "lookup_" + "service_metric_" * 3000}, + }, + "max_tokens": 10, + } + assert _approximate_input_size(body) >= TOKENIZE_OFF_EVENT_LOOP_MIN_CHARS + + _, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth(token="key-tool-choice", spend=0.0, max_budget=100.0) + threads = [] + + def recording_token_counter(**kwargs): + threads.append(threading.current_thread()) + return 1000 + + with patch.object(litellm, "token_counter", recording_token_counter): + reservation = await reserve_budget_for_request( + request_body=body, + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert threads + assert all(thread is not threading.main_thread() for thread in threads) + + +@pytest.mark.asyncio +async def test_small_prompt_is_tokenized_inline(spend_counter_state): + """A thread hand-off costs more than counting a small prompt""" + _, key_cache = spend_counter_state + proxy_logging_obj = ProxyLogging(user_api_key_cache=key_cache) + valid_token = UserAPIKeyAuth(token="key-inline", spend=0.0, max_budget=100.0) + threads = [] + + def recording_token_counter(**kwargs): + threads.append(threading.current_thread()) + return 10 + + with patch.object(litellm, "token_counter", recording_token_counter): + reservation = await reserve_budget_for_request( + request_body=_request_body(), + route="/chat/completions", + llm_router=None, + valid_token=valid_token, + team_object=None, + user_object=None, + prisma_client=None, + user_api_key_cache=key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + assert reservation is not None + assert threads == [threading.main_thread()] diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index 355c6d27eb2..510fb977a61 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -3,7 +3,7 @@ import copy import datetime import json from types import SimpleNamespace -from typing import AsyncGenerator, Callable, Optional +from typing import AsyncGenerator, Callable, Final, Optional from unittest.mock import AsyncMock, MagicMock, patch import httpx diff --git a/tests/test_litellm/proxy/test_component_allowlists.py b/tests/test_litellm/proxy/test_component_allowlists.py index 3dd8d8b28cd..0fdb43d60da 100644 --- a/tests/test_litellm/proxy/test_component_allowlists.py +++ b/tests/test_litellm/proxy/test_component_allowlists.py @@ -75,6 +75,7 @@ _DB_ENV_KEYS = ( "DATABASE_HOST_READ_REPLICA", "DATABASE_PASSWORD", "IAM_TOKEN_DB_AUTH", + "AZURE_POSTGRESQL_AUTH", ) _PRE_DB_ENV = {_key: os.environ.pop(_key, None) for _key in _DB_ENV_KEYS} _PRE_COMPONENT_LIFESPAN = app.router.lifespan_context diff --git a/tests/test_litellm/proxy/test_prisma_migration.py b/tests/test_litellm/proxy/test_prisma_migration.py new file mode 100644 index 00000000000..01b768ea8dc --- /dev/null +++ b/tests/test_litellm/proxy/test_prisma_migration.py @@ -0,0 +1,68 @@ +import os +from unittest.mock import MagicMock, patch + +import pytest + +from litellm.proxy import prisma_migration + + +class TestPrismaMigration: + @patch("litellm.proxy.prisma_migration.subprocess.run") + @patch("litellm.proxy.prisma_migration.run_server") + def test_main_enforces_migration_check_by_default( + self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock + ) -> None: + mock_subprocess_run.return_value = MagicMock(returncode=0, stdout="", stderr="") + + with patch.dict(os.environ, {}, clear=True): + assert prisma_migration.main() == 0 + + mock_run_server.assert_called_once_with( + ("--skip_server_startup", "--enforce_prisma_migration_check"), + standalone_mode=False, + ) + + @patch("litellm.proxy.prisma_migration.subprocess.run") + @patch("litellm.proxy.prisma_migration.run_server") + def test_main_disables_migration_check_when_explicitly_false( + self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock + ) -> None: + mock_subprocess_run.return_value = MagicMock(returncode=0, stdout="", stderr="") + + with patch.dict(os.environ, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}, clear=True): + assert prisma_migration.main() == 0 + + mock_run_server.assert_called_once_with(("--skip_server_startup",), standalone_mode=False) + + @patch("litellm.proxy.prisma_migration.subprocess.run") + @patch("litellm.proxy.prisma_migration.run_server") + def test_main_returns_prisma_generate_exit_code_when_enforced( + self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock + ) -> None: + mock_subprocess_run.return_value = MagicMock(returncode=7, stdout="", stderr="") + + with patch.dict(os.environ, {}, clear=True): + assert prisma_migration.main() == 7 + + @patch("litellm.proxy.prisma_migration.subprocess.run") + @patch("litellm.proxy.prisma_migration.run_server") + def test_main_ignores_prisma_generate_exit_code_when_disabled( + self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock + ) -> None: + mock_subprocess_run.return_value = MagicMock(returncode=7, stdout="", stderr="") + + with patch.dict(os.environ, {"ENFORCE_PRISMA_MIGRATION_CHECK": "false"}, clear=True): + assert prisma_migration.main() == 0 + + @patch("litellm.proxy.prisma_migration.subprocess.run") + @patch("litellm.proxy.prisma_migration.run_server") + def test_main_propagates_migration_failure( + self, mock_run_server: MagicMock, mock_subprocess_run: MagicMock + ) -> None: + mock_run_server.side_effect = SystemExit(1) + + with patch.dict(os.environ, {}, clear=True): + with pytest.raises(SystemExit, match="1"): + prisma_migration.main() + + mock_subprocess_run.assert_not_called() diff --git a/tests/test_litellm/proxy/test_proxy_cli.py b/tests/test_litellm/proxy/test_proxy_cli.py index 20d17b5a510..28f43345350 100644 --- a/tests/test_litellm/proxy/test_proxy_cli.py +++ b/tests/test_litellm/proxy/test_proxy_cli.py @@ -1585,6 +1585,7 @@ class TestProxyInitializationHelpers: "DATABASE_URL": "", "DIRECT_URL": "", "IAM_TOKEN_DB_AUTH": "", + "AZURE_POSTGRESQL_AUTH": "", "USE_AWS_KMS": "", } with patch.dict(os.environ, env_overrides): @@ -2335,3 +2336,85 @@ class TestPostgresStatementTimeoutOptions: standalone_mode=False, ) return {k: os.environ[k] for k in ("DATABASE_URL", "DIRECT_URL") if k in os.environ} + + +class TestTokenAuthCliFlags: + """`--azure_postgresql_auth` has to reach the URL assembly the same way the env var does.""" + + def _invoke_with_azure_host(self, args): + from click.testing import CliRunner + + from litellm.proxy.db.token_auth import build_azure_entra_token_provider + from litellm.proxy.proxy_cli import run_server + + build_azure_entra_token_provider.cache_clear() + clean_env = { + k: v + for k, v in os.environ.items() + if k + not in ( + "DATABASE_URL", + "DIRECT_URL", + "IAM_TOKEN_DB_AUTH", + "AZURE_POSTGRESQL_AUTH", + "DATABASE_URL_READ_REPLICA", + ) + } + clean_env["DATABASE_HOST"] = "writer.postgres.database.azure.com" + clean_env["DATABASE_USER"] = "litellm@contoso.onmicrosoft.com" + clean_env["DATABASE_NAME"] = "litellm_db" + + mock_proxy_module = MagicMock( + app=MagicMock(), + ProxyConfig=MagicMock(), + KeyManagementSettings=MagicMock(), + save_worker_config=MagicMock(), + ) + with ( + patch.dict(os.environ, clean_env, clear=True), + patch.dict( + "sys.modules", + { + "proxy_server": mock_proxy_module, + "litellm.proxy.proxy_server": mock_proxy_module, + }, + ), + patch( + "litellm.secret_managers.get_azure_ad_token_provider.get_azure_ad_token_provider", + return_value=lambda: "ENTRA_TOKEN", + ), + patch("litellm.proxy.db.prisma_client.should_update_prisma_schema", return_value=False), + patch("litellm.proxy.db.prisma_client.PrismaManager.setup_database"), + patch("uvicorn.run"), + patch( + "litellm.proxy.proxy_cli.ProxyInitializationHelpers._get_default_unvicorn_init_args" + ) as mock_get_args, + ): + mock_get_args.return_value = { + "app": "litellm.proxy.proxy_server:app", + "host": "localhost", + "port": 8000, + } + result = CliRunner().invoke(run_server, args) + database_url = os.getenv("DATABASE_URL") + toggle = os.getenv("AZURE_POSTGRESQL_AUTH") + build_azure_entra_token_provider.cache_clear() + return result, database_url, toggle + + def test_azure_flag_assembles_a_token_bearing_database_url(self): + result, database_url, toggle = self._invoke_with_azure_host( + ["--local", "--azure_postgresql_auth"] + ) + + assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}" + assert database_url is not None + assert "ENTRA_TOKEN" in database_url + assert "writer.postgres.database.azure.com" in database_url + assert toggle == "True" + + def test_without_the_flag_no_token_is_minted(self): + result, database_url, toggle = self._invoke_with_azure_host(["--local"]) + + assert result.exit_code == 0, f"exit_code={result.exit_code}, output={result.output}" + assert "ENTRA_TOKEN" not in (database_url or "") + assert toggle is None diff --git a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py index 849d5494c6e..772b08bc9d0 100644 --- a/tests/test_litellm/proxy/test_redis_auth_cache_flag.py +++ b/tests/test_litellm/proxy/test_redis_auth_cache_flag.py @@ -29,7 +29,7 @@ class _FakeRedisCache(RedisCache): network calls are made. """ - def __init__(self): # noqa: super().__init__ skipped intentionally + def __init__(self): # super().__init__ skipped intentionally self._store = {} def set_cache(self, key, value, **kwargs): # type: ignore[override] diff --git a/tests/test_litellm/proxy/test_spend_log_cleanup.py b/tests/test_litellm/proxy/test_spend_log_cleanup.py index 87fbdd4c933..ce0b6b755cc 100644 --- a/tests/test_litellm/proxy/test_spend_log_cleanup.py +++ b/tests/test_litellm/proxy/test_spend_log_cleanup.py @@ -793,6 +793,7 @@ async def test_spend_logs_retention_alone_does_not_touch_the_session_rollup(): tables = [call[0][0] for call in client.db.execute_raw.call_args_list] assert any('"LiteLLM_SpendLogs"' in sql for sql in tables) assert not any('"LiteLLM_AutoRouterSession"' in sql for sql in tables) + assert not any('"LiteLLM_HealthCheckTable"' in sql for sql in tables) @pytest.mark.asyncio @@ -807,25 +808,47 @@ async def test_session_retention_alone_cleans_only_the_session_rollup(): @pytest.mark.asyncio -async def test_each_retention_key_cuts_off_at_its_own_horizon(): - from datetime import datetime, timezone +async def test_health_check_retention_alone_cleans_only_the_health_check_table(): + client = _mock_prisma_for_retention([0]) + cleaner = SpendLogCleanup(general_settings={"maximum_health_check_retention_period": "30d"}) + cleaner.pod_lock_manager = None + await cleaner.cleanup_old_spend_logs(client) + tables = [call[0][0] for call in client.db.execute_raw.call_args_list] + assert len(tables) == 1 + assert '"LiteLLM_HealthCheckTable"' in tables[0] + assert '"health_check_id"' in tables[0] + assert '"checked_at"' in tables[0] + cutoff_date = client.db.execute_raw.call_args[0][1] + expected_cutoff = datetime.now(timezone.utc) - timedelta(days=30) + assert abs((cutoff_date - expected_cutoff).total_seconds()) < 1 - client = _mock_prisma_for_retention([0, 0, 0]) + +@pytest.mark.asyncio +async def test_each_retention_key_cuts_off_at_its_own_horizon(): + client = _mock_prisma_for_retention([0, 0, 0, 0]) cleaner = SpendLogCleanup( general_settings={ "maximum_spend_logs_retention_period": "7d", "maximum_autorouter_session_retention_period": "365d", + "maximum_health_check_retention_period": "30d", } ) cleaner.pod_lock_manager = None await cleaner.cleanup_old_spend_logs(client) cutoffs = { - ("LiteLLM_AutoRouterSession" if '"LiteLLM_AutoRouterSession"' in call[0][0] else "logs"): call[0][1] + ( + "LiteLLM_AutoRouterSession" + if '"LiteLLM_AutoRouterSession"' in call[0][0] + else "LiteLLM_HealthCheckTable" + if '"LiteLLM_HealthCheckTable"' in call[0][0] + else "logs" + ): call[0][1] for call in client.db.execute_raw.call_args_list } now = datetime.now(timezone.utc) assert (now - cutoffs["logs"]).days == 7 assert (now - cutoffs["LiteLLM_AutoRouterSession"]).days == 365 + assert (now - cutoffs["LiteLLM_HealthCheckTable"]).days == 30 @pytest.mark.asyncio @@ -910,6 +933,32 @@ async def test_run_budget_is_shared_across_tables_not_granted_per_table(): assert "LiteLLM_SpendLogs" in tables_touched +@pytest.mark.asyncio +async def test_cleanup_groups_share_budget_so_health_checks_still_get_a_delete(): + mock_prisma_client = MagicMock() + _wire_tx(mock_prisma_client.db) + mock_db = MagicMock() + _wire_tx(mock_db) + mock_db.execute_raw = AsyncMock(return_value=1000) + mock_prisma_client.db = mock_db + + cleaner = SpendLogCleanup( + general_settings={ + "maximum_spend_logs_retention_period": "7d", + "maximum_health_check_retention_period": "30d", + "maximum_spend_logs_cleanup_max_batches": 50, + "maximum_spend_logs_cleanup_run_budget": "1s", + } + ) + cleaner.pod_lock_manager = None + + await cleaner.cleanup_old_spend_logs(mock_prisma_client) + + tables_touched = {call[0][0].split('"')[1] for call in mock_db.execute_raw.call_args_list} + assert "LiteLLM_SpendLogs" in tables_touched + assert "LiteLLM_HealthCheckTable" in tables_touched + + @pytest.mark.asyncio async def test_each_batch_carries_a_statement_and_lock_timeout(): """ diff --git a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py index 30fd4a74bb0..18b02ac7772 100644 --- a/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py +++ b/tests/test_litellm/proxy/utils/prisma_and_spend/test_prisma_client_lifecycle.py @@ -36,7 +36,7 @@ async def test_prismaclient_init_wires_default_config( proxy_logging_obj=proxy_logging, ) pinned = { - "iam_token_db_auth": pc.iam_token_db_auth, + "token_auth": pc.token_auth, "db_reconnect_cooldown_seconds": pc._db_reconnect_cooldown_seconds, "db_health_watchdog_interval_seconds": pc._db_health_watchdog_interval_seconds, "db_health_watchdog_enabled": pc._db_health_watchdog_enabled, @@ -48,7 +48,7 @@ async def test_prismaclient_init_wires_default_config( "db_reconnect_lock_is_lock": isinstance(pc._db_reconnect_lock, asyncio.Lock), } assert pinned == { - "iam_token_db_auth": None, + "token_auth": None, "db_reconnect_cooldown_seconds": 15, "db_health_watchdog_interval_seconds": 30, "db_health_watchdog_enabled": True, diff --git a/tests/test_litellm/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py b/tests/test_litellm/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py index 6752d76847f..bff6f261020 100644 --- a/tests/test_litellm/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py +++ b/tests/test_litellm/router_utils/pre_call_checks/test_prompt_caching_deployment_check.py @@ -1,3 +1,5 @@ +import asyncio +import copy import os import sys from typing import List, cast @@ -9,6 +11,8 @@ sys.path.insert(0, os.path.abspath("../..")) import litellm from litellm.caching.dual_cache import DualCache from litellm.constants import DEFAULT_MINIMUM_PROMPT_CACHE_TOKEN_COUNT +from litellm.integrations.anthropic_cache_control_hook import AnthropicCacheControlHook +from litellm.integrations.custom_logger import CustomLogger from litellm.router_utils.pre_call_checks.prompt_caching_deployment_check import ( PromptCachingDeploymentCheck, _get_min_token_count_for_deployments, @@ -187,6 +191,210 @@ async def test_async_filter_deployments_narrows_for_group_whose_model_minimum_is assert filtered == [deployments[1]] +AUTO_CACHING_MODEL = "anthropic/claude-sonnet-4-5" + + +def _auto_caching_messages() -> List[AllMessageValues]: + """A prompt over the model minimum that carries no client cache_control.""" + return cast( + List[AllMessageValues], + [ + {"role": "system", "content": "word " * 3000}, + {"role": "user", "content": "hello"}, + ], + ) + + +def _affinity_messages(messages: List[AllMessageValues]) -> List[AllMessageValues]: + """The messages the check keys deployment affinity on, for a group of `AUTO_CACHING_MODEL`.""" + return AnthropicCacheControlHook.messages_with_default_injections( + messages=messages, + models=(AUTO_CACHING_MODEL,), + ) + + +class _SentMessagesCapture(CustomLogger): + def __init__(self): + self.messages: List[AllMessageValues] | None = None + + async def async_log_success_event(self, kwargs, response_obj, start_time, end_time): + standard_logging_object = kwargs.get("standard_logging_object") + if standard_logging_object is not None: + self.messages = standard_logging_object["messages"] + + +async def _eventually(predicate, timeout: float = 10.0): + """Success callbacks run as tasks, so give the write a bounded window to land.""" + deadline = asyncio.get_running_loop().time() + timeout + while asyncio.get_running_loop().time() < deadline: + result = predicate() + if result: + return result + await asyncio.sleep(0.05) + return predicate() + + +@pytest.mark.asyncio +async def test_affinity_key_matches_the_messages_auto_caching_actually_sends(monkeypatch, local_model_cost_map): + """ + The regression. `enable_anthropic_prompt_caching` injects cache_control inside + `litellm.acompletion`, which runs after routing, so at filter time the messages carried no + marker, `extract_cacheable_prefix` returned [], the key was None, and the check no-opped on + every request. Routing must derive the same key the success event writes from the messages the + request was actually sent with, otherwise auto-injected caching gets no affinity at all. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + capture = _SentMessagesCapture() + monkeypatch.setattr(litellm, "callbacks", [capture]) + messages = _auto_caching_messages() + + await litellm.acompletion( + model=AUTO_CACHING_MODEL, + messages=copy.deepcopy(messages), + mock_response="ok", + api_key="sk-fake", + ) + sent_messages = await _eventually(lambda: capture.messages) + assert sent_messages is not None + + routing_key = PromptCachingCache.get_prompt_caching_cache_key(_affinity_messages(messages), None) + + assert routing_key is not None + assert routing_key == PromptCachingCache.get_prompt_caching_cache_key(sent_messages, None) + + +@pytest.mark.asyncio +async def test_repeated_auto_cached_prefix_pins_to_one_deployment(monkeypatch, local_model_cost_map): + """ + End to end over the router: identical requests with no client cache_control must stop bouncing + across a multi-deployment group once one deployment has cached the prefix. Bedrock and Anthropic + caches are per account and region, so every bounce paid the cache write premium and never read. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + router = litellm.Router( + model_list=[ + { + "model_name": MODEL_GROUP_ALIAS, + "litellm_params": {"model": AUTO_CACHING_MODEL, "api_key": "sk-fake"}, + "model_info": {"id": model_id}, + } + for model_id in ("dep-1", "dep-2") + ], + optional_pre_call_checks=["prompt_caching"], + ) + messages = _auto_caching_messages() + + first = await router.acompletion(model=MODEL_GROUP_ALIAS, messages=messages, mock_response="ok") + served_by = first._hidden_params["model_id"] + + affinity_key = PromptCachingCache.get_prompt_caching_cache_key(_affinity_messages(messages), None) + assert await _eventually(lambda: router.cache.get_cache(key=affinity_key)) is not None + + subsequent = [ + (await router.acompletion(model=MODEL_GROUP_ALIAS, messages=messages, mock_response="ok"))._hidden_params[ + "model_id" + ] + for _ in range(4) + ] + + assert subsequent == [served_by] * 4 + + +@pytest.mark.asyncio +async def test_per_request_enable_prompt_caching_reaches_the_affinity_key(monkeypatch, local_model_cost_map): + """ + `enable_prompt_caching` turns auto-injection on for a single request while the global flag stays + off, so routing has to read it too. Ignore it and the key comes off unmarked messages, which is + never what the request goes on to send, and the pin is lost for every per-key enablement. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", False) + cache = DualCache() + check = PromptCachingDeploymentCheck(cache=cache) + deployments = _deployments(AUTO_CACHING_MODEL, AUTO_CACHING_MODEL) + messages = _auto_caching_messages() + + sent = AnthropicCacheControlHook.messages_with_default_injections( + messages=messages, models=(AUTO_CACHING_MODEL,), enable_prompt_caching=True + ) + assert sent != messages + await PromptCachingCache(cache=cache).async_add_model_id(model_id="dep-2", messages=sent, tools=None) + + filtered = await check.async_filter_deployments( + model=MODEL_GROUP_ALIAS, + healthy_deployments=deployments, + messages=messages, + request_kwargs={"enable_prompt_caching": True}, + ) + + assert filtered == [deployments[1]] + + +@pytest.mark.asyncio +async def test_tool_marked_cache_control_keeps_routing_off_another_requests_prefix(monkeypatch, local_model_cost_map): + """ + Tools carrying the client's own cache_control make auto-injection stand down, so this request + will not carry litellm's breakpoints. Routing must see the tools as well. Ignore them and it + keys off the injected prefix, pinning the request to whichever deployment cached a different, + tool-less request whose prefix it can never actually reuse. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + cache = DualCache() + check = PromptCachingDeploymentCheck(cache=cache) + deployments = _deployments(AUTO_CACHING_MODEL, AUTO_CACHING_MODEL) + messages = _auto_caching_messages() + cache_marked_tools = [ + { + "type": "function", + "function": {"name": "get_weather", "parameters": {"type": "object", "properties": {}}}, + "cache_control": {"type": "ephemeral"}, + } + ] + + await PromptCachingCache(cache=cache).async_add_model_id( + model_id="dep-2", messages=_affinity_messages(messages), tools=None + ) + + without_tools = await check.async_filter_deployments( + model=MODEL_GROUP_ALIAS, healthy_deployments=deployments, messages=messages + ) + assert without_tools == [deployments[1]] + + with_tools = await check.async_filter_deployments( + model=MODEL_GROUP_ALIAS, + healthy_deployments=deployments, + messages=messages, + request_kwargs={"tools": cache_marked_tools}, + ) + + assert with_tools == deployments + + +def test_client_supplied_cache_control_keeps_its_own_prefix_boundary(monkeypatch, local_model_cost_map): + """ + Auto-injection stands down when the client marks its own breakpoints, so the affinity key must + keep keying off the client's boundary. Injecting on top would push the boundary to the trailing + turn and break affinity for prompts that already worked. + """ + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + messages = cast( + List[AllMessageValues], + [ + { + "role": "system", + "content": [ + {"type": "text", "text": "word " * 3000, "cache_control": {"type": "ephemeral"}}, + ], + }, + {"role": "user", "content": "hello"}, + ], + ) + + for_key = _affinity_messages(messages) + + assert for_key is messages + assert PromptCachingCache.extract_cacheable_prefix(for_key) == messages[:1] + + @pytest.mark.asyncio async def test_wildcard_route_resolves_underlying_model_minimum(local_model_cost_map): from litellm import Router diff --git a/tests/test_litellm/rust_bridge/test_chat_completions.py b/tests/test_litellm/rust_bridge/test_chat_completions.py new file mode 100644 index 00000000000..47cb66932b7 --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_chat_completions.py @@ -0,0 +1,420 @@ +"""Tests for the Rust chat completions bridge. + +The native callables are dependency-injected through +``set_rust_chat_completions`` rather than patched, so these run without the +compiled extension present. +""" + +from __future__ import annotations + +import pytest + +import litellm +from litellm.rust_bridge import chat_completions as bridge +from litellm.types.utils import ModelResponse + +RUST_RESPONSE = { + "created": 1_700_000_000, + "model": "claude-sonnet-4-5-20260101", + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "hello from rust"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": 11, + "completion_tokens": 4, + "total_tokens": 15, + "prompt_tokens_details": { + "cached_tokens": 0, + "cache_creation_tokens": 0, + "text_tokens": 11, + }, + }, +} + +MESSAGES = [{"role": "user", "content": "hi"}] + + +class _FakeDeclined(Exception): + """Stands in for the native `RustBridgeDeclined`.""" + + +class _FakeUpstream(Exception): + """Stands in for the native `RustUpstreamError`; args are (status, message).""" + + +class _FakeNative: + RustBridgeDeclined = _FakeDeclined + RustUpstreamError = _FakeUpstream + + +def _fake_native_bridge(monkeypatch): + """Expose the bridge's exception classes without the compiled extension.""" + monkeypatch.setattr(bridge, "get_native_bridge", lambda: _FakeNative()) + + +def _hide_native_bridge(monkeypatch): + """Simulate a wheel built without the compiled extension. + + There is no injection seam for "the .so is absent", so the loader itself is + replaced; every other case here uses `set_rust_chat_completions`. + """ + monkeypatch.setattr(bridge, "get_native_bridge", lambda: None) + + +@pytest.fixture(autouse=True) +def reset_bridge(): + """Every test starts with no injected callables, and leaves none behind.""" + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) + yield + bridge.set_rust_chat_completions( + chat_completions=None, achat_completions=None, decline=None + ) + + +class _RecordingDecline: + """A stand-in for the native gate that records what it was asked.""" + + def __init__(self, reason: str | None = None): + self.reason = reason + self.calls: list[dict] = [] + + def __call__(self, **kwargs): + self.calls.append(kwargs) + return self.reason + + +class _RecordingCall: + def __init__(self, result=None, error: Exception | None = None): + self.result = result if result is not None else dict(RUST_RESPONSE) + self.error = error + self.calls: list[dict] = [] + + def __call__(self, **kwargs): + self.calls.append(kwargs) + if self.error is not None: + raise self.error + return self.result + + +class _RecordingAsyncCall(_RecordingCall): + async def __call__(self, **kwargs): + return _RecordingCall.__call__(self, **kwargs) + + +def _accepts(**overrides) -> bool: + kwargs = { + "model": "claude-sonnet-4-5", + "messages": MESSAGES, + "optional_params": {"max_tokens": 16}, + "custom_llm_provider": "anthropic", + "litellm_params": {"rust": True}, + "stream": None, + } + kwargs.update(overrides) + return bridge.rust_chat_completions_accepts(**kwargs) + + +class TestGate: + def test_declines_when_the_deployment_did_not_opt_in(self, monkeypatch): + monkeypatch.delenv("LITELLM_RUST", raising=False) + gate = _RecordingDecline() + bridge.set_rust_chat_completions(decline=gate) + assert _accepts(litellm_params={}) is False + assert _accepts(litellm_params=None) is False + assert _accepts(litellm_params={"rust": False}) is False + assert gate.calls == [], "the gate must not be consulted before opt-in" + + def test_accepts_when_the_deployment_opted_in_and_the_core_agrees(self, monkeypatch): + monkeypatch.delenv("LITELLM_RUST", raising=False) + gate = _RecordingDecline() + bridge.set_rust_chat_completions(decline=gate) + assert _accepts() is True + assert gate.calls[0]["model"] == "claude-sonnet-4-5" + assert gate.calls[0]["custom_llm_provider"] == "anthropic" + + def test_the_env_var_opts_in_without_a_per_model_flag(self, monkeypatch): + monkeypatch.setenv("LITELLM_RUST", "true") + bridge.set_rust_chat_completions(decline=_RecordingDecline()) + assert _accepts(litellm_params={}) is True + + def test_declines_streaming_and_providers_off_the_path(self, monkeypatch): + monkeypatch.delenv("LITELLM_RUST", raising=False) + gate = _RecordingDecline() + bridge.set_rust_chat_completions(decline=gate) + assert _accepts(stream=True) is False + assert _accepts(custom_llm_provider="openai") is False + assert _accepts(custom_llm_provider=None) is False + assert gate.calls == [] + + def test_declines_an_anthropic_request_carrying_a_litellm_metadata_user_id(self, monkeypatch): + """`AnthropicConfig.transform_request` copies a valid `user_id` into the Messages body. + + It does that inside the function the Rust route replaces, and the core is + handed `optional_params` only, so accepting here would send the request + to Anthropic with the abuse-detection attribution silently missing. + """ + monkeypatch.delenv("LITELLM_RUST", raising=False) + gate = _RecordingDecline() + bridge.set_rust_chat_completions(decline=gate) + assert _accepts(litellm_params={"rust": True, "metadata": {"user_id": "u-123"}}) is False + assert gate.calls == [], "the core must not be consulted for a request it cannot see the key of" + + # Bedrock's Converse transform reads no `user_id`, and an Anthropic request + # whose metadata carries none is one Python would not attribute either. + assert ( + _accepts( + custom_llm_provider="bedrock", + model="bedrock/us-east-1/anthropic.claude-v2", + litellm_params={"rust": True, "metadata": {"user_id": "u-123"}}, + ) + is True + ) + assert _accepts(litellm_params={"rust": True, "metadata": {"trace_id": "t-1"}}) is True + assert _accepts(litellm_params={"rust": True, "metadata": {"user_id": None}}) is True + assert _accepts(litellm_params={"rust": True, "metadata": None}) is True + + def test_declines_a_bedrock_request_while_the_proxy_owns_request_metadata(self, monkeypatch): + """`AmazonConverseConfig` resolves proxy-owned `requestMetadata` onto the + Converse body from `litellm_params`, and owning that field also means + evicting a caller-supplied one. The core can do neither, so an operator + who armed `bedrock_request_metadata_fields` keeps the Python path. + """ + monkeypatch.delenv("LITELLM_RUST", raising=False) + gate = _RecordingDecline() + bridge.set_rust_chat_completions(decline=gate) + bedrock = { + "custom_llm_provider": "bedrock", + "model": "bedrock/us-east-1/anthropic.claude-v2", + } + + monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", ["user_api_key_team_id"]) + assert _accepts(**bedrock) is False + assert gate.calls == [], "the core must not be consulted for a field it cannot write" + assert _accepts() is True, "arming Bedrock attribution must not decline Anthropic" + + monkeypatch.setattr(litellm, "bedrock_request_metadata_fields", None) + assert _accepts(**bedrock) is True, "the decline follows the operator's opt-in alone" + + def test_declines_when_the_core_declines(self, monkeypatch): + monkeypatch.delenv("LITELLM_RUST", raising=False) + bridge.set_rust_chat_completions(decline=_RecordingDecline("streaming")) + assert _accepts() is False + + def test_declines_when_the_bridge_is_unavailable(self, monkeypatch): + monkeypatch.delenv("LITELLM_RUST", raising=False) + _hide_native_bridge(monkeypatch) + assert _accepts() is False + + def test_declines_when_the_gate_itself_raises(self, monkeypatch): + monkeypatch.delenv("LITELLM_RUST", raising=False) + + def exploding(**_kwargs): + raise RuntimeError("boom") + + bridge.set_rust_chat_completions(decline=exploding) + assert _accepts() is False + + +def _call_kwargs(model_response: ModelResponse) -> dict: + return { + "model": "claude-sonnet-4-5", + "messages": MESSAGES, + "optional_params": {"max_tokens": 16}, + "model_response": model_response, + "api_key": "sk-test", + "api_base": None, + "custom_llm_provider": "anthropic", + "extra_headers": {}, + "timeout": 30.0, + "on_response": lambda _rust_response: None, + } + + +class TestSyncCall: + def test_builds_a_model_response_and_stamps_the_rust_header(self): + native = _RecordingCall() + bridge.set_rust_chat_completions(chat_completions=native) + model_response = ModelResponse() + original_id = model_response.id + + result = bridge.chat_completions(**_call_kwargs(model_response)) + + assert result is not None + assert result.choices[0].message.content == "hello from rust" + assert result.choices[0].finish_reason == "stop" + assert result.model == "claude-sonnet-4-5-20260101" + assert result.usage.prompt_tokens == 11 + assert result.usage.completion_tokens == 4 + assert result.usage.total_tokens == 15 + assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + assert result.id == original_id, ( + "the rust path must keep the chatcmpl id litellm already minted" + ) + + def test_passes_the_timeout_through_as_seconds(self): + native = _RecordingCall() + bridge.set_rust_chat_completions(chat_completions=native) + bridge.chat_completions(**_call_kwargs(ModelResponse())) + assert native.calls[0]["timeout_seconds"] == 30.0 + + def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): + _hide_native_bridge(monkeypatch) + assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + + def test_falls_back_when_the_core_declines_before_calling_the_provider(self, monkeypatch): + _fake_native_bridge(monkeypatch) + bridge.set_rust_chat_completions( + chat_completions=_RecordingCall(error=_FakeDeclined("streaming")) + ) + assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + + +class TestAsyncCall: + @pytest.mark.asyncio + async def test_builds_a_model_response(self): + bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) + result = await bridge.achat_completions(**_call_kwargs(ModelResponse())) + assert result is not None + assert result.choices[0].message.content == "hello from rust" + assert result._hidden_params["additional_headers"] == {"x-litellm-rust": "true"} + + @pytest.mark.asyncio + async def test_falls_back_when_the_bridge_is_unavailable(self, monkeypatch): + _hide_native_bridge(monkeypatch) + assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None + + @pytest.mark.asyncio + async def test_falls_back_when_the_core_declines_before_calling_the_provider( + self, monkeypatch + ): + _fake_native_bridge(monkeypatch) + bridge.set_rust_chat_completions( + achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")) + ) + assert await bridge.achat_completions(**_call_kwargs(ModelResponse())) is None + + +class TestAsyncFallbackWrapper: + @pytest.mark.asyncio + async def test_returns_the_rust_response_without_running_the_fallback(self): + bridge.set_rust_chat_completions(achat_completions=_RecordingAsyncCall()) + ran = [] + + async def fallback(): + ran.append(True) + return "python" + + result = await bridge.achat_completions_or_fallback( + **_call_kwargs(ModelResponse()), python_fallback=fallback + ) + assert result.choices[0].message.content == "hello from rust" + assert ran == [] + + @pytest.mark.asyncio + async def test_runs_the_fallback_when_the_core_declines(self, monkeypatch): + _fake_native_bridge(monkeypatch) + bridge.set_rust_chat_completions( + achat_completions=_RecordingAsyncCall(error=_FakeDeclined("streaming")) + ) + + async def fallback(): + return "python" + + result = await bridge.achat_completions_or_fallback( + **_call_kwargs(ModelResponse()), python_fallback=fallback + ) + assert result == "python" + + @pytest.mark.asyncio + async def test_runs_the_fallback_when_the_bridge_is_unavailable(self, monkeypatch): + _hide_native_bridge(monkeypatch) + + async def fallback(): + return "python" + + result = await bridge.achat_completions_or_fallback( + **_call_kwargs(ModelResponse()), python_fallback=fallback + ) + assert result == "python" + + +class TestFailureClassification: + """A failure the provider already saw must not be retried on the Python + path: it would bill the customer for the same work twice.""" + + @pytest.fixture(autouse=True) + def _native_exceptions(self, monkeypatch): + _fake_native_bridge(monkeypatch) + + def test_a_decline_falls_back_because_nothing_was_sent(self): + bridge.set_rust_chat_completions( + chat_completions=_RecordingCall(error=_FakeDeclined("streaming")) + ) + assert bridge.chat_completions(**_call_kwargs(ModelResponse())) is None + + def test_an_upstream_failure_is_surfaced_with_its_status(self): + from litellm.exceptions import APIError + + bridge.set_rust_chat_completions( + chat_completions=_RecordingCall(error=_FakeUpstream(429, "429: rate limited")) + ) + with pytest.raises(APIError) as raised: + bridge.chat_completions(**_call_kwargs(ModelResponse())) + assert raised.value.status_code == 429 + assert "rate limited" in str(raised.value) + + def test_a_transport_failure_with_no_response_surfaces_as_a_500(self): + from litellm.exceptions import APIError + + bridge.set_rust_chat_completions( + chat_completions=_RecordingCall(error=_FakeUpstream(0, "connection reset")) + ) + with pytest.raises(APIError) as raised: + bridge.chat_completions(**_call_kwargs(ModelResponse())) + assert raised.value.status_code == 500 + + def test_an_unrecognized_error_is_not_swallowed(self): + bridge.set_rust_chat_completions( + chat_completions=_RecordingCall(error=RuntimeError("something else")) + ) + with pytest.raises(RuntimeError): + bridge.chat_completions(**_call_kwargs(ModelResponse())) + + @pytest.mark.asyncio + async def test_the_async_wrapper_does_not_fall_back_on_an_upstream_failure(self): + from litellm.exceptions import APIError + + bridge.set_rust_chat_completions( + achat_completions=_RecordingAsyncCall(error=_FakeUpstream(500, "500: boom")) + ) + ran = [] + + async def fallback(): + ran.append(True) + return "python" + + with pytest.raises(APIError): + await bridge.achat_completions_or_fallback( + **_call_kwargs(ModelResponse()), python_fallback=fallback + ) + assert ran == [], "a request the provider already served must not be re-issued" + + @pytest.mark.asyncio + async def test_the_async_wrapper_falls_back_on_a_decline(self): + bridge.set_rust_chat_completions( + achat_completions=_RecordingAsyncCall(error=_FakeDeclined("blank message text")) + ) + + async def fallback(): + return "python" + + result = await bridge.achat_completions_or_fallback( + **_call_kwargs(ModelResponse()), python_fallback=fallback + ) + assert result == "python" diff --git a/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py b/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py index f02f59cccc0..cee0da79802 100644 --- a/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py +++ b/tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py @@ -11,6 +11,10 @@ import pytest from litellm.secret_managers.get_azure_ad_token_provider import ( get_azure_ad_token_provider, + infer_credential_type_from_environment, +) +from litellm.types.secret_managers.get_azure_ad_token_provider import ( + AzureCredentialType, ) @@ -215,6 +219,46 @@ class TestGetAzureAdTokenProvider: token = result() assert token == "mock-certificate-token" + @patch.dict( + os.environ, + { + "AZURE_CLIENT_ID": "test-client-id", + "AZURE_TENANT_ID": "test-tenant-id", + "AZURE_FEDERATED_TOKEN_FILE": "/var/run/secrets/azure/tokens/azure-identity-token", + "AZURE_AUTHORITY_HOST": "https://login.microsoftonline.com/", + }, + clear=True, + ) + @patch("azure.identity.get_bearer_token_provider") + @patch("azure.identity.ManagedIdentityCredential") + @patch("azure.identity.DefaultAzureCredential") + def test_get_azure_ad_token_provider_prefers_workload_identity_over_managed_identity( + self, + mock_default_azure_credential, + mock_managed_identity_credential, + mock_get_bearer_token_provider, + ): + """The AKS workload identity webhook injects AZURE_CLIENT_ID, AZURE_TENANT_ID, and + AZURE_FEDERATED_TOKEN_FILE, and never a client secret. Reading the bare client id as a + managed identity sends the pod to IMDS, which has no identity attached to it, so every + token request fails and the federated token is never exchanged. Only + DefaultAzureCredential's chain reaches WorkloadIdentityCredential.""" + mock_credential_instance = MagicMock() + mock_default_azure_credential.return_value = mock_credential_instance + mock_get_bearer_token_provider.return_value = MagicMock( + return_value="mock-workload-identity-token" + ) + + result = get_azure_ad_token_provider() + + assert ( + infer_credential_type_from_environment() + == AzureCredentialType.DefaultAzureCredential + ) + mock_managed_identity_credential.assert_not_called() + mock_default_azure_credential.assert_called_once_with() + assert result() == "mock-workload-identity-token" + @patch.dict(os.environ, {}, clear=True) # Clear all environment variables @patch("azure.identity.get_bearer_token_provider") @patch("azure.identity.DefaultAzureCredential") diff --git a/tests/test_litellm/test_assert_ci_coverage.py b/tests/test_litellm/test_assert_ci_coverage.py index 59cfff52992..4c8f4cc2984 100644 --- a/tests/test_litellm/test_assert_ci_coverage.py +++ b/tests/test_litellm/test_assert_ci_coverage.py @@ -209,3 +209,74 @@ def test_character_class_globs_match_the_letter_shards_circleci_uses(): def test_the_repo_as_it_stands_has_no_unrecorded_slice_gap(): findings = coverage._deselected_everywhere(coverage._load_allowlist()) assert [f.subject for f in findings] == [] + + +def _allowlist(*, tests: tuple[str, ...] = (), dockerfiles: tuple[str, ...] = ()): + return coverage.Allowlist( + test_paths=(coverage.AllowEntry(paths=tests, reason="r"),) if tests else (), + dockerfiles=(coverage.AllowEntry(paths=dockerfiles, reason="r"),) if dockerfiles else (), + ) + + +def test_an_allowlist_path_whose_file_is_gone_is_reported(): + findings = coverage._stale_allowlist_paths( + _allowlist(tests=("tests/gone/test_a.py",)), + test_files=("tests/live/test_b.py",), + dockerfiles=(), + ) + assert [f.subject for f in findings] == ["tests/gone/test_a.py"] + assert "test_paths" in findings[0].detail + + +def test_an_allowlist_path_that_still_matches_a_file_is_left_alone(): + findings = coverage._stale_allowlist_paths( + _allowlist(tests=("tests/live/test_b.py",)), + test_files=("tests/live/test_b.py",), + dockerfiles=(), + ) + assert findings == () + + +def test_a_directory_entry_survives_while_any_file_below_it_remains(): + findings = coverage._stale_allowlist_paths( + _allowlist(tests=("tests/live",)), + test_files=("tests/live/nested/test_b.py",), + dockerfiles=(), + ) + assert findings == () + + +def test_a_glob_entry_matching_nothing_is_reported_like_any_other(): + findings = coverage._stale_allowlist_paths( + _allowlist(tests=("tests/live/test_z*.py",)), + test_files=("tests/live/test_b.py",), + dockerfiles=(), + ) + assert [f.subject for f in findings] == ["tests/live/test_z*.py"] + + +def test_a_stale_dockerfile_entry_is_named_under_its_own_section(): + findings = coverage._stale_allowlist_paths( + _allowlist(dockerfiles=("docker/Dockerfile.gone",)), + test_files=(), + dockerfiles=("docker/Dockerfile.database",), + ) + assert [(f.subject, "dockerfiles" in f.detail) for f in findings] == [("docker/Dockerfile.gone", True)] + + +def test_the_repo_as_it_stands_has_no_stale_allowlist_entry(): + findings = coverage._stale_allowlist_paths( + coverage._load_allowlist(), + test_files=coverage._test_files(), + dockerfiles=coverage._dockerfiles(), + ) + assert [f.subject for f in findings] == [] + + +def test_a_dockerfile_directory_entry_is_stale_because_only_an_exact_path_exempts_one(): + findings = coverage._stale_allowlist_paths( + _allowlist(dockerfiles=("docker",)), + test_files=(), + dockerfiles=("docker/Dockerfile.database",), + ) + assert [f.subject for f in findings] == ["docker"] diff --git a/tests/test_litellm/test_assert_workflow_dir_hygiene.py b/tests/test_litellm/test_assert_workflow_dir_hygiene.py new file mode 100644 index 00000000000..37f3b4518bb --- /dev/null +++ b/tests/test_litellm/test_assert_workflow_dir_hygiene.py @@ -0,0 +1,115 @@ +"""Tests for .github/scripts/assert_workflow_dir_hygiene.py.""" + +import importlib.util +import sys +from pathlib import Path +from typing import Final + +import pytest + +_REPO_ROOT: Final = Path(__file__).resolve().parents[2] +_MODULE_PATH: Final = _REPO_ROOT / ".github" / "scripts" / "assert_workflow_dir_hygiene.py" +_spec: Final = importlib.util.spec_from_file_location("assert_workflow_dir_hygiene", _MODULE_PATH) +hygiene: Final = importlib.util.module_from_spec(_spec) +sys.modules[_spec.name] = hygiene # @dataclass(slots=True) rebuilds via sys.modules +_spec.loader.exec_module(hygiene) + + +def _codes(path_name, triggers): + return [f.code for f in hygiene._naming_findings(Path(path_name), frozenset(triggers))] + + +def test_a_call_only_workflow_without_the_prefix_is_flagged(): + assert _codes("deploy.yml", {"workflow_call"}) == ["WF002"] + + +def test_a_call_only_workflow_with_the_prefix_is_clean(): + assert _codes("_deploy.yml", {"workflow_call"}) == [] + + +def test_a_dual_mode_workflow_keeps_its_plain_name(): + # workflow_call plus a human trigger is deliberate: the `_` prefix would hide a + # workflow someone is meant to be able to dispatch. + assert _codes("create-release-branch.yml", {"workflow_call", "workflow_dispatch"}) == [] + + +def test_a_prefixed_workflow_nobody_can_call_is_flagged(): + assert _codes("_helper.yml", {"push"}) == ["WF003"] + + +def test_a_plain_workflow_with_ordinary_triggers_is_clean(): + assert _codes("test-unit.yml", {"pull_request", "push"}) == [] + + +@pytest.mark.parametrize( + "raw, expected", + [ + ({"on": "push"}, {"push"}), + ({"on": ["push", "pull_request"]}, {"push", "pull_request"}), + ({"on": {"workflow_call": None}}, {"workflow_call"}), + ({True: {"pull_request": None}}, {"pull_request"}), + ({"jobs": {}}, set()), + ("not a mapping", set()), + ], +) +def test_triggers_reads_every_shape_the_on_key_takes(raw, expected): + # YAML 1.1 turns a bare `on:` key into the boolean True, which is why the loaded + # document has to be read both ways. + assert hygiene._triggers(raw) == frozenset(expected) + + +def test_the_repo_as_it_stands_holds_only_workflows_in_the_workflow_dir(): + assert [f.subject for f in hygiene._strays(hygiene.WORKFLOW_DIR)] == [] + + +def test_the_repo_as_it_stands_names_every_reusable_workflow_with_the_prefix(): + assert [f.subject for f in hygiene._misnamed(hygiene.WORKFLOW_DIR)] == [] + + +def test_the_repo_as_it_stands_spells_every_workflow_yml(): + assert [f.subject for f in hygiene._misspelled(hygiene.WORKFLOW_DIR)] == [] + + +_WORKFLOW: Final = "name: ci\non: [push]\njobs:\n a:\n runs-on: ubuntu-latest\n steps: [{run: 'true'}]\n" + + +def _populate(directory, files): + for name, body in files.items(): + target = directory / name + target.parent.mkdir(parents=True, exist_ok=True) + target.write_text(body, encoding="utf-8") + return directory + + +def _findings(directory): + return [ + (f.subject, f.code) + for f in hygiene._strays(directory) + hygiene._misspelled(directory) + hygiene._misnamed(directory) + ] + + +def test_a_script_at_the_top_level_is_a_stray(tmp_path): + directory = _populate(tmp_path, {"ci.yml": _WORKFLOW, "render.py": "print(1)\n"}) + assert _findings(directory) == [("render.py", "WF001")] + + +def test_a_script_inside_a_subdirectory_is_left_alone(tmp_path): + directory = _populate(tmp_path, {"ci.yml": _WORKFLOW, "helpers/render.py": "print(1)\n"}) + assert _findings(directory) == [] + + +def test_a_yaml_workflow_is_a_naming_finding_not_a_stray(tmp_path): + directory = _populate(tmp_path, {"test-model-map.yaml": _WORKFLOW}) + assert _findings(directory) == [("test-model-map.yaml", "WF004")] + + +def test_the_yaml_message_names_the_rename_and_not_the_scripts_directory(tmp_path): + directory = _populate(tmp_path, {"test-model-map.yaml": _WORKFLOW}) + detail = hygiene._misspelled(directory)[0].detail + assert "test-model-map.yml" in detail + assert hygiene.SCRIPT_HOME not in detail + + +def test_a_yaml_workflow_is_still_held_to_the_prefix_rules(tmp_path): + directory = _populate(tmp_path, {"deploy.yaml": "on: {workflow_call: null}\njobs: {}\n"}) + assert _findings(directory) == [("deploy.yaml", "WF004"), ("deploy.yaml", "WF002")] diff --git a/tests/test_litellm/test_check_test_quality.py b/tests/test_litellm/test_check_test_quality.py index 7f8ce4c36d0..4fea5761cc8 100644 --- a/tests/test_litellm/test_check_test_quality.py +++ b/tests/test_litellm/test_check_test_quality.py @@ -397,3 +397,154 @@ def test_a_none_comparison_reads_as_absence(tmp_path): def test_a_membership_test_without_the_negation_is_left_alone(tmp_path): source = _MEMBERSHIP_GATE.replace('"ACME_API_KEY" not in os.environ', '"ACME_API_KEY" in os.environ') assert _codes(tmp_path, source) == [] + + +_SNAPSHOT_CONFTEST = """import litellm +import pytest + + +@pytest.fixture(autouse=True) +def restore_globals(): + original_state = {} + original_state["drop_params"] = litellm.drop_params + for attr in ("api_base", "num_retries"): + original_state[attr] = getattr(litellm, attr) + yield + for attr, value in original_state.items(): + setattr(litellm, attr, value) +""" + + +def _conftest_codes(tmp_path, source, name="conftest.py"): + snippet = tmp_path / name + snippet.write_text(source, encoding="utf-8") + return [v.code for v in checker.check_file(snippet)] + + +def test_every_snapshotted_global_is_counted_once(tmp_path): + assert _conftest_codes(tmp_path, _SNAPSHOT_CONFTEST) == ["TQ007", "TQ007", "TQ007"] + + +def test_the_names_come_from_the_loop_tuple_as_well_as_the_direct_keys(tmp_path): + snippet = tmp_path / "conftest.py" + snippet.write_text(_SNAPSHOT_CONFTEST, encoding="utf-8") + reported = [v.message.split("`")[1] for v in checker.check_file(snippet)] + assert sorted(reported) == ["litellm.api_base", "litellm.drop_params", "litellm.num_retries"] + + +def test_the_same_global_saved_twice_counts_once(tmp_path): + source = _SNAPSHOT_CONFTEST.replace( + '("api_base", "num_retries")', '("api_base", "num_retries", "drop_params")' + ) + assert _conftest_codes(tmp_path, source) == ["TQ007", "TQ007", "TQ007"] + + +def test_the_rule_only_looks_at_conftest_files(tmp_path): + assert _conftest_codes(tmp_path, _SNAPSHOT_CONFTEST, name="test_snapshot.py") == [] + + +def test_a_conftest_that_snapshots_nothing_is_clean(tmp_path): + source = "import pytest\n\n\n@pytest.fixture\ndef client():\n return object()\n" + assert _conftest_codes(tmp_path, source) == [] + + +def test_a_snapshot_entry_is_suppressible_with_a_reason(tmp_path): + source = _SNAPSHOT_CONFTEST.replace( + 'original_state["drop_params"] = litellm.drop_params', + 'original_state["drop_params"] = litellm.drop_params # test-quality-ok: owned by the SDK config surface', + ) + assert _conftest_codes(tmp_path, source) == ["TQ007", "TQ007"] + + +_NAMED_MAPPING_CONFTEST = """import litellm +import pytest + +_SCALAR_DEFAULTS = { + "num_retries": None, + "set_verbose": False, +} +_EXTRA_ATTRS = ("api_base", "drop_params") + + +@pytest.fixture(autouse=True) +def restore_globals(): + original_state = {} + for attr in _SCALAR_DEFAULTS: + original_state[attr] = getattr(litellm, attr) + for attr in _EXTRA_ATTRS: + original_state[attr] = getattr(litellm, attr) + yield + for attr, value in original_state.items(): + setattr(litellm, attr, value) +""" + + +def test_a_save_loop_over_a_module_level_dict_counts_its_keys(tmp_path): + # The two largest inventories in the repo name their list instead of spelling it + # out, so a rule that only reads literal iterables sees neither. + reported = [v.message.split("`")[1] for v in checker.check_file(_written(tmp_path, _NAMED_MAPPING_CONFTEST))] + assert sorted(reported) == [ + "litellm.api_base", + "litellm.drop_params", + "litellm.num_retries", + "litellm.set_verbose", + ] + + +def test_a_named_iterable_that_is_not_a_module_constant_is_skipped_quietly(tmp_path): + source = _NAMED_MAPPING_CONFTEST.replace("for attr in _EXTRA_ATTRS:", "for attr in dir(litellm):") + reported = [v.message.split("`")[1] for v in checker.check_file(_written(tmp_path, source))] + assert sorted(reported) == ["litellm.num_retries", "litellm.set_verbose"] + + +def _written(tmp_path, source, name="conftest.py"): + path = tmp_path / name + path.write_text(source, encoding="utf-8") + return path + + +_HELPER_DICT_CONFTEST = """import litellm +import pytest + +_CALLBACK_ATTRS = ("callbacks", "success_callback") + + +def _copy_litellm_state(): + state = {} + for attr in _CALLBACK_ATTRS: + if hasattr(litellm, attr): + value = getattr(litellm, attr) + state[attr] = value.copy() if isinstance(value, list) else value + return state + + +@pytest.fixture(autouse=True) +def restore_globals(): + saved = _copy_litellm_state() + yield + for attr, value in saved.items(): + setattr(litellm, attr, value) +""" + + +def test_a_snapshot_built_in_a_helper_under_any_dict_name_is_counted(tmp_path): + # Two conftests build their inventory inside a helper and call the dict `state`, + # so a rule keyed on blessed dict names sees neither. + reported = [v.message.split("`")[1] for v in checker.check_file(_written(tmp_path, _HELPER_DICT_CONFTEST))] + assert sorted(reported) == ["litellm.callbacks", "litellm.success_callback"] + + +def test_the_read_may_sit_a_statement_above_the_store(tmp_path): + # `val = getattr(litellm, attr)` then `state[attr] = val.copy()` is the common + # shape; requiring the store itself to read litellm loses every one of them. + source = _HELPER_DICT_CONFTEST.replace( + " state[attr] = value.copy() if isinstance(value, list) else value", + " state[attr] = list(value)", + ) + reported = [v.message.split("`")[1] for v in checker.check_file(_written(tmp_path, source))] + assert sorted(reported) == ["litellm.callbacks", "litellm.success_callback"] + + +def test_a_loop_storing_under_a_key_that_is_not_the_loop_variable_is_not_an_inventory(tmp_path): + source = _HELPER_DICT_CONFTEST.replace("state[attr] =", 'state["fixed"] =') + assert [v.code for v in checker.check_file(_written(tmp_path, source))] == [] diff --git a/tests/test_litellm/test_logging.py b/tests/test_litellm/test_logging.py index 784ec5b6cf4..8551085cbd6 100644 --- a/tests/test_litellm/test_logging.py +++ b/tests/test_litellm/test_logging.py @@ -2,15 +2,14 @@ import ast import asyncio import json import os +import re import sys from pathlib import Path from typing import List import pytest -sys.path.insert( - 0, os.path.abspath("../../..") -) # Adds the parent directory to the system-path +sys.path.insert(0, os.path.abspath("../../..")) # Adds the parent directory to the system-path import logging import sys @@ -20,7 +19,10 @@ from litellm._logging import ( CorrelationContextFilter, CorrelationPlainFormatter, JsonFormatter, + SecretRedactionFilter, + StdoutLogTruncationFilter, _initialize_loggers_with_handler, + _stdout_truncation_marker, _turn_on_json, session_id_var, set_session_id, @@ -30,6 +32,7 @@ from litellm._logging import ( verbose_proxy_logger, verbose_router_logger, ) +from litellm.constants import LITELLM_TRUNCATED_PAYLOAD_FIELD from litellm.integrations.custom_logger import CustomLogger from litellm.types.utils import StandardLoggingPayload @@ -238,9 +241,7 @@ def test_json_formatter_includes_component_field(): ) output = formatter.format(record) obj = json.loads(output) - assert ( - obj["component"] == logger_name - ), f"Expected component={logger_name!r}, got {obj.get('component')!r}" + assert obj["component"] == logger_name, f"Expected component={logger_name!r}, got {obj.get('component')!r}" def test_json_formatter_includes_logger_field(): @@ -260,9 +261,7 @@ def test_json_formatter_includes_logger_field(): ) output = formatter.format(record) obj = json.loads(output) - assert ( - obj["logger"] == "proxy_server.py:123" - ), f"Expected logger='proxy_server.py:123', got {obj['logger']!r}" + assert obj["logger"] == "proxy_server.py:123", f"Expected logger='proxy_server.py:123', got {obj['logger']!r}" def test_json_formatter_extra_component_not_overwritten(): @@ -281,9 +280,7 @@ def test_json_formatter_extra_component_not_overwritten(): ) record.component = "auth-service" obj = json.loads(formatter.format(record)) - assert ( - obj["component"] == "auth-service" - ), f"User-supplied component was overwritten, got {obj['component']!r}" + assert obj["component"] == "auth-service", f"User-supplied component was overwritten, got {obj['component']!r}" def test_initialize_loggers_with_handler_sets_propagate_false(): @@ -295,9 +292,9 @@ def test_initialize_loggers_with_handler_sets_propagate_false(): # Check that propagate is set to False for all loggers for logger in ALL_LOGGERS: - assert ( - logger.propagate is False - ), f"Logger {logger.name} has propagate set to {logger.propagate}, expected False" + assert logger.propagate is False, ( + f"Logger {logger.name} has propagate set to {logger.propagate}, expected False" + ) @pytest.mark.asyncio @@ -335,9 +332,9 @@ async def test_cache_hit_includes_custom_llm_provider(): await asyncio.sleep(0.5) # Verify we have logged events - assert ( - len(test_custom_logger.logged_standard_logging_payloads) >= 2 - ), f"Expected at least 2 logged events, got {len(test_custom_logger.logged_standard_logging_payloads)}" + assert len(test_custom_logger.logged_standard_logging_payloads) >= 2, ( + f"Expected at least 2 logged events, got {len(test_custom_logger.logged_standard_logging_payloads)}" + ) # Find the cache hit event (should be the second call) cache_hit_payload = None @@ -347,20 +344,18 @@ async def test_cache_hit_includes_custom_llm_provider(): break # Verify cache hit event was found - assert ( - cache_hit_payload is not None - ), "No cache hit event found in logged payloads" + assert cache_hit_payload is not None, "No cache hit event found in logged payloads" # Verify custom_llm_provider is included in the cache hit payload - assert ( - "custom_llm_provider" in cache_hit_payload - ), "custom_llm_provider missing from cache hit standard logging payload" + assert "custom_llm_provider" in cache_hit_payload, ( + "custom_llm_provider missing from cache hit standard logging payload" + ) # Verify custom_llm_provider has a valid value (should be "openai" for gpt-3.5-turbo) custom_llm_provider = cache_hit_payload["custom_llm_provider"] - assert ( - custom_llm_provider is not None and custom_llm_provider != "" - ), f"custom_llm_provider should not be None or empty, got: {custom_llm_provider}" + assert custom_llm_provider is not None and custom_llm_provider != "", ( + f"custom_llm_provider should not be None or empty, got: {custom_llm_provider}" + ) print( f"Cache hit standard logging payload with custom_llm_provider: {custom_llm_provider}", @@ -666,6 +661,171 @@ def test_set_trace_id_strips_control_characters(): trace_id_var.reset(token) +_MARKER_RE = re.compile(rf"\.\.\. \({LITELLM_TRUNCATED_PAYLOAD_FIELD} skipped (\d+) chars\..*?\) \.\.\.", re.S) + + +def _extract_marker(text: str) -> "re.Match[str] | None": + return _MARKER_RE.search(text) + + +def _make_record(level: int, msg: str, args=(), exc_info=None) -> logging.LogRecord: + return logging.LogRecord( + name="LiteLLM Router", + level=level, + pathname="", + lineno=0, + msg=msg, + args=args, + exc_info=exc_info, + ) + + +def test_oversized_info_record_is_truncated(monkeypatch): + """An error string echoing a huge request payload must not reach stdout in full.""" + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "500") + payload = "p" * 100_000 + record = _make_record(logging.INFO, "litellm.acompletion(model=%s) Exception %s", ("gpt-4", payload)) + + assert StdoutLogTruncationFilter().filter(record) is True + + message = record.getMessage() + assert LITELLM_TRUNCATED_PAYLOAD_FIELD in message + assert len(message) <= 500 + assert message.startswith("litellm.acompletion(model=gpt-4) Exception ppp") + assert message.endswith("ppp") + + marker = _extract_marker(message) + assert marker is not None + kept, skipped = len(message) - len(marker.group(0)), int(marker.group(1)) + assert kept + skipped == 43 + len(payload) + + +def test_truncated_message_fits_the_configured_cap(monkeypatch): + """The cap is the whole point of the setting, so the marker has to be paid for out of + the budget instead of appended on top of a limit-sized head and tail.""" + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "500") + record = _make_record(logging.ERROR, "Exception %s", ("p" * 2000,)) + + assert StdoutLogTruncationFilter().filter(record) is True + + message = record.getMessage() + assert _extract_marker(message) is not None + assert len(message) == 500 + + +@pytest.mark.parametrize("payload_len", [501, 512, 1000, 9999, 100_000]) +def test_truncated_message_never_exceeds_the_cap(monkeypatch, payload_len): + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "500") + record = _make_record(logging.ERROR, "%s", ("p" * payload_len,)) + + assert StdoutLogTruncationFilter().filter(record) is True + + assert len(record.getMessage()) <= 500 + + +_NO_BUDGET_PAYLOAD = "p" * 2000 +_MARKER_SIZED_CAP = len(_stdout_truncation_marker(len(_NO_BUDGET_PAYLOAD))) + + +@pytest.mark.parametrize("cap", [_MARKER_SIZED_CAP, _MARKER_SIZED_CAP - 1, 100]) +def test_cap_leaving_no_room_for_the_marker_still_bounds_output(monkeypatch, cap): + """An operator can set the cap at or below the marker's own length, leaving nothing to + spend on a head and tail, and the output still has to fit.""" + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", str(cap)) + record = _make_record(logging.ERROR, "%s", (_NO_BUDGET_PAYLOAD,)) + + assert StdoutLogTruncationFilter().filter(record) is True + + assert len(record.getMessage()) == cap + + +def test_debug_record_is_not_truncated(monkeypatch): + """--detailed_debug exists to dump full payloads, so DEBUG records pass through.""" + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "500") + payload = "p" * 100_000 + record = _make_record(logging.DEBUG, "raw request %s", (payload,)) + + assert StdoutLogTruncationFilter().filter(record) is True + + assert record.getMessage() == f"raw request {payload}" + + +def test_truncation_disabled_by_zero_limit(monkeypatch): + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "0") + payload = "p" * 100_000 + record = _make_record(logging.ERROR, "Exception %s", (payload,)) + + assert StdoutLogTruncationFilter().filter(record) is True + + assert record.getMessage() == f"Exception {payload}" + + +def test_oversized_traceback_is_truncated(monkeypatch): + """verbose_proxy_logger.exception() re-logs the payload inside the traceback too.""" + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "500") + try: + raise ValueError("payload " + "p" * 100_000) + except ValueError: + exc_info = sys.exc_info() + record = _make_record(logging.ERROR, "Exception occured", exc_info=exc_info) + + assert StdoutLogTruncationFilter().filter(record) is True + + assert record.exc_text is not None + assert LITELLM_TRUNCATED_PAYLOAD_FIELD in record.exc_text + assert len(record.exc_text) <= 500 + assert "Traceback (most recent call last)" in record.exc_text + + +def test_falsy_exc_info_is_not_formatted(monkeypatch): + """Callers pass exc_info=False, which logging leaves on the record as a bool.""" + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "500") + record = _make_record(logging.WARNING, "skipping malformed endpoint %s", ("p" * 100_000,), exc_info=False) + + assert StdoutLogTruncationFilter().filter(record) is True + + assert record.exc_text is None + assert LITELLM_TRUNCATED_PAYLOAD_FIELD in record.getMessage() + + +def test_secret_filter_keeps_truncated_traceback(monkeypatch): + """SecretRedactionFilter runs after truncation, so it must redact the capped + traceback instead of reformatting the full one from exc_info.""" + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "500") + try: + raise ValueError("sk-1234567890abcdefghij payload " + "p" * 100_000) + except ValueError: + exc_info = sys.exc_info() + record = _make_record(logging.ERROR, "Exception occured", exc_info=exc_info) + + assert StdoutLogTruncationFilter().filter(record) is True + assert SecretRedactionFilter().filter(record) is True + + assert record.exc_text is not None + assert len(record.exc_text) <= 500 + assert "sk-1234567890abcdefghij" not in record.exc_text + + +def test_truncation_filter_survives_json_reconfiguration(): + """The cap lives on the loggers, so swapping handlers (JSON mode) can't drop it.""" + _turn_on_json() + + for lg in (verbose_logger, verbose_router_logger, verbose_proxy_logger): + assert any(isinstance(f, StdoutLogTruncationFilter) for f in lg.filters), f"{lg.name} lost stdout truncation" + + +def test_oversized_error_is_truncated_end_to_end(monkeypatch, caplog): + """The router's own exception log line must come out bounded, not just the filter in isolation.""" + monkeypatch.setenv("MAX_STRING_LENGTH_STDOUT_LOG", "500") + + with caplog.at_level(logging.INFO, logger="LiteLLM Router"): + verbose_router_logger.info("litellm.acompletion(model=%s) Exception %s", "gpt-4", "p" * 100_000) + + emitted = "".join(record.getMessage() for record in caplog.records) + assert LITELLM_TRUNCATED_PAYLOAD_FIELD in emitted + assert len(emitted) <= 500 + + def test_set_session_id_bounds_length(): """set_session_id() must bound length so an oversized caller-supplied value isn't repeated across every log line for the request.""" @@ -674,4 +834,3 @@ def test_set_session_id_bounds_length(): assert len(session_id_var.get()) == 256 finally: session_id_var.reset(token) - diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 65debae9a16..4a06a5dfb2e 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -4,6 +4,7 @@ import json import logging import os import sys +import threading from unittest.mock import AsyncMock, MagicMock, patch import pytest @@ -3344,6 +3345,277 @@ def test_pre_call_checks_counts_once_and_filters_on_max_input_tokens(monkeypatch assert calls == [1] +def test_pre_call_checks_uses_precounted_tokens(monkeypatch): + """ + An async caller counts off the event loop and passes the result in. _pre_call_checks + must filter on that count instead of re-counting on the loop. + """ + router = litellm.Router( + model_list=[ + {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, + ], + enable_pre_call_checks=True, + ) + monkeypatch.setattr( + router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} + ) + + calls = [] + monkeypatch.setattr( + litellm, "token_counter", lambda *a, **k: calls.append(1) or 1 + ) + + deployments = [ + {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, + ] + with pytest.raises(litellm.ContextWindowExceededError): + router._pre_call_checks( + model="m", + healthy_deployments=deployments, + messages=[{"role": "user", "content": "hi"}], + input_token_count=1000, + ) + + assert calls == [] + + +async def test_async_get_healthy_deployments_counts_tokens_off_the_event_loop(monkeypatch): + """ + The async deployment path must hand _pre_call_checks a count taken in a worker thread, + so a multi-MB prompt never blocks the proxy during deployment selection. + """ + router = litellm.Router( + model_list=[ + {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, + ], + enable_pre_call_checks=True, + ) + monkeypatch.setattr( + router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 1_000_000} + ) + + counting_threads = [] + monkeypatch.setattr( + litellm, + "token_counter", + lambda *a, **k: counting_threads.append(threading.current_thread()) or 42, + ) + + counts_passed_in = [] + original_pre_call_checks = router._pre_call_checks + + def spy(**kwargs): + counts_passed_in.append(kwargs.get("input_token_count")) + return original_pre_call_checks(**kwargs) + + monkeypatch.setattr(router, "_pre_call_checks", spy) + + result = await router.async_get_healthy_deployments( + model="m", + request_kwargs={}, + messages=[{"role": "user", "content": "hi"}], + input=None, + specific_deployment=False, + parent_otel_span=None, + ) + + assert len(result) == 1 + assert counts_passed_in == [42] + assert len(counting_threads) == 1 + assert counting_threads[0] is not threading.current_thread() + + +@pytest.mark.parametrize( + "model_info,expected", + [ + ({"max_input_tokens": 100}, True), + ({"max_input_tokens": None}, False), + ({}, False), + ], +) +def test_pre_call_checks_need_token_count(monkeypatch, model_info, expected): + """Only a deployment that declares an integer context window makes a token count worth taking.""" + router = litellm.Router( + model_list=[ + {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, + ], + enable_pre_call_checks=True, + ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: model_info) + + deployments = [ + {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, + ] + assert router._pre_call_checks_need_token_count("m", deployments) is expected + + +def test_deployment_max_input_tokens_survives_an_unmappable_deployment(monkeypatch): + """ + _pre_call_checks skips a deployment it cannot resolve and carries on. The off-loop + pre-count must do the same, or an unmapped first deployment hides the limit declared by + a later one and the count lands back on the event loop. + """ + router = litellm.Router( + model_list=[ + {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, + ], + enable_pre_call_checks=True, + ) + + def flaky_model_info(deployment, received_model_name, id=None): + if deployment["model_info"]["id"] == "unmapped": + raise ValueError("This model isn't mapped yet.") + return {"max_input_tokens": 100} + + monkeypatch.setattr(router, "get_router_model_info", flaky_model_info) + + unmapped = {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "unmapped"}} + mapped = {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "mapped"}} + + assert router._deployment_max_input_tokens("m", unmapped) is None + assert router._deployment_max_input_tokens("m", mapped) == 100 + assert router._pre_call_checks_need_token_count("m", [unmapped, mapped]) is True + + +def test_pre_call_checks_does_not_recount_inline_after_an_off_loop_failure(monkeypatch): + """ + When the off-loop count failed there is nothing left to filter on, so _pre_call_checks must + return the deployments unfiltered rather than repeating the count on the event loop. + """ + router = litellm.Router( + model_list=[ + {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, + ], + enable_pre_call_checks=True, + ) + monkeypatch.setattr( + router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} + ) + + calls = [] + monkeypatch.setattr( + litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 + ) + + deployments = [ + {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, + ] + result = router._pre_call_checks( + model="m", + healthy_deployments=deployments, + messages=[{"role": "user", "content": "hi"}], + input_token_count=None, + skip_inline_token_count=True, + ) + + assert calls == [] + assert len(result) == 1 + + +async def test_async_get_healthy_deployments_never_recounts_on_the_loop(monkeypatch): + """ + An off-loop count that raises must not send the same work back onto the event loop through + _pre_call_checks' inline fallback. + """ + router = litellm.Router( + model_list=[ + {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, + ], + enable_pre_call_checks=True, + ) + monkeypatch.setattr( + router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} + ) + + counting_threads = [] + + def exploding_counter(*args, **kwargs): + counting_threads.append(threading.current_thread()) + raise ValueError("Invalid content item type: image") + + monkeypatch.setattr(litellm, "token_counter", exploding_counter) + + result = await router.async_get_healthy_deployments( + model="m", + request_kwargs={}, + messages=[{"role": "user", "content": "hi"}], + input=None, + specific_deployment=False, + parent_otel_span=None, + ) + + assert len(result) == 1 + assert len(counting_threads) == 1 + assert counting_threads[0] is not threading.current_thread() + + +async def test_acount_pre_call_check_tokens_leaves_the_event_loop_free(monkeypatch): + """ + A multi-MB prompt must not stall the proxy: a competing coroutine has to get + scheduled while the router's context-window count is in flight. + """ + router = litellm.Router( + model_list=[ + {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, + ], + enable_pre_call_checks=True, + ) + monkeypatch.setattr( + router, "get_router_model_info", lambda **kwargs: {"max_input_tokens": 5} + ) + + deployments = [ + {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, + ] + ran = [] + + async def competitor(): + ran.append("competitor") + + task = asyncio.create_task(competitor()) + count = await router._acount_pre_call_check_tokens( + model="m", + healthy_deployments=deployments, + messages=[{"role": "user", "content": "A" * 512 * 1024}], + input=None, + request_kwargs=None, + ) + ran.append("count") + await task + + assert count is not None and count > 0 + assert ran == ["competitor", "count"] + + +async def test_acount_pre_call_check_tokens_skips_without_max_input_tokens(monkeypatch): + """No deployment limits its context window, so there is nothing to count.""" + router = litellm.Router( + model_list=[ + {"model_name": "m", "litellm_params": {"model": "gpt-3.5-turbo"}}, + ], + enable_pre_call_checks=True, + ) + monkeypatch.setattr(router, "get_router_model_info", lambda **kwargs: {}) + + calls = [] + monkeypatch.setattr( + litellm, "token_counter", lambda *a, **k: calls.append(1) or 1000 + ) + + count = await router._acount_pre_call_check_tokens( + model="m", + healthy_deployments=[ + {"litellm_params": {"model": "gpt-3.5-turbo"}, "model_info": {"id": "d1"}}, + ], + messages=[{"role": "user", "content": "hi"}], + input=None, + request_kwargs=None, + ) + + assert count is None + assert calls == [] + + def test_pre_call_checks_counts_tokens_from_responses_input_string(monkeypatch): """ Responses API calls pass `input` (str) instead of `messages`. Context-window diff --git a/tests/test_litellm/test_router_model_cost_isolation.py b/tests/test_litellm/test_router_model_cost_isolation.py index 4674a8b1dfa..5bb854c12e0 100644 --- a/tests/test_litellm/test_router_model_cost_isolation.py +++ b/tests/test_litellm/test_router_model_cost_isolation.py @@ -20,6 +20,7 @@ sys.path.insert( import litellm from litellm import Router +from litellm.litellm_core_utils.ptu_pricing import ptu_config_error from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo from litellm.utils import ( _invalidate_model_cost_lowercase_map, @@ -42,6 +43,16 @@ def _simulate_price_data_reload(fetched_catalog): reapply_runtime_model_cost_registrations() +def _nested_container_ids(value: object) -> frozenset[int]: + """Identities of every dict/list reachable from `value`, so two structures can be + checked for shared mutable state without writing into either one.""" + if isinstance(value, dict): + return frozenset({id(value)} | {i for v in value.values() for i in _nested_container_ids(v)}) + if isinstance(value, list): + return frozenset({id(value)} | {i for v in value for i in _nested_container_ids(v)}) + return frozenset() + + def _restore_model_cost_entries(original_entries): for key, value in original_entries.items(): if value is None: @@ -1689,11 +1700,234 @@ def test_nothing_is_zeroed_while_the_feature_is_off(): @pytest.mark.parametrize("dropped", ["team_id", "ptu_effective_from"], ids=["no team_id", "no ptu_effective_from"]) -def test_a_deployment_the_rollup_will_not_charge_is_not_zeroed(dropped): - """The rollup refuses to price a reservation missing either field, so zeroing on the - looser count-and-rate test alone would leave the deployment serving for free with - nothing charged in its place.""" +def test_an_incomplete_reservation_is_refused_rather_than_served(dropped): + """POST /model/new answers 400 for exactly this config, so config.yaml must not quietly + accept it. Serving it would bill per token while accruing no flat cost, which is the + state the operator was trying to leave.""" incomplete = {k: v for k, v in _PTU_MODEL_INFO.items() if k != dropped} - entry = _ptu_router(model_info=incomplete, litellm_params={"input_cost_per_token": 5e-06}).model_list[0] + + with pytest.raises(ValueError) as raised: + _ptu_router(model_info=incomplete, litellm_params={"input_cost_per_token": 5e-06}) + + assert "gpt-4o-ptu" in str(raised.value) + + +@pytest.mark.parametrize( + "dropped, expected", + [ + ("team_id", "team_id is required when PTU fields are set (one model maps to one team)"), + ("cost_per_ptu_per_hour", "ptu_count and cost_per_ptu_per_hour must be set together"), + ], + ids=["no team_id", "count without rate"], +) +def test_the_refusal_reason_is_the_one_the_model_endpoint_answers_with(dropped, expected): + """One rule, stated once. If these drift, an operator gets contradictory guidance + depending on which path they used.""" + incomplete = {k: v for k, v in _PTU_MODEL_INFO.items() if k != dropped} + + assert ptu_config_error(incomplete) == expected + with pytest.raises(ValueError) as raised: + _ptu_router(model_info=incomplete) + + assert expected in str(raised.value) + + +@pytest.mark.parametrize("dropped", ["team_id", "ptu_effective_from"], ids=["no team_id", "no ptu_effective_from"]) +def test_an_incomplete_reservation_is_left_alone_while_the_feature_is_off(dropped): + """Nothing accrues with the flag off, so refusing a deployment there would take a + serving model away from an operator who never opted in.""" + incomplete = {k: v for k, v in _PTU_MODEL_INFO.items() if k != dropped} + entry = _ptu_router( + model_info=incomplete, litellm_params={"input_cost_per_token": 5e-06}, ptu_enabled=False + ).model_list[0] assert entry["litellm_params"]["input_cost_per_token"] == 5e-06 + + +@pytest.mark.parametrize("dropped", ["team_id", "ptu_effective_from"], ids=["no team_id", "no ptu_effective_from"]) +def test_the_proxy_drops_the_deployment_rather_than_failing_to_boot(dropped): + """The proxy builds its router with ignore_invalid_deployments, so one bad entry must + cost that entry and not the whole config.""" + incomplete = {k: v for k, v in _PTU_MODEL_INFO.items() if k != dropped} + with patch.dict(os.environ, {"LITELLM_ENABLE_PTU_COST_ATTRIBUTION": "True"}, clear=False): + router = Router( + model_list=[ + { + "model_name": "gpt-4o-ptu", + "litellm_params": {"model": "anthropic/claude-sonnet-4-5-20250929", "api_key": "sk-not-used"}, + "model_info": dict(incomplete), + }, + { + "model_name": "plain-sibling", + "litellm_params": {"model": "anthropic/claude-sonnet-4-5-20250929", "api_key": "sk-not-used"}, + }, + ], + ignore_invalid_deployments=True, + ) + + assert [entry["model_name"] for entry in router.model_list] == ["plain-sibling"] + + +def test_a_complete_reservation_still_registers(): + """The refusal must be scoped to a broken reservation, not to PTU configuration.""" + entry = _ptu_router().model_list[0] + + assert entry["model_name"] == "gpt-4o-ptu" + assert entry["litellm_params"]["input_cost_per_token"] == 0.0 + + +def test_nested_custom_model_info_does_not_pollute_shared_backend(): + backend_model = "gpt-4o-search-preview" + custom_id = "lit5471-search-custom" + sibling_id = "lit5471-search-sibling" + builtin_info = copy.deepcopy(litellm.get_model_info(model=backend_model)) + expected_nested = copy.deepcopy(builtin_info["search_context_cost_per_query"]) + model_keys = { + backend_model: copy.deepcopy(litellm.model_cost.get(backend_model)), + custom_id: copy.deepcopy(litellm.model_cost.get(custom_id)), + sibling_id: copy.deepcopy(litellm.model_cost.get(sibling_id)), + } + try: + router = Router( + model_list=[ + { + "model_name": "search-custom", + "litellm_params": {"model": backend_model, "api_key": "fake-key"}, + "model_info": { + "id": custom_id, + "search_context_cost_per_query": { + "search_context_size_low": 0.123, + }, + }, + }, + { + "model_name": "search-sibling", + "litellm_params": {"model": backend_model, "api_key": "fake-key"}, + "model_info": {"id": sibling_id}, + }, + ], + ) + + custom_info = router.get_deployment_model_info(model_id=custom_id, model_name=backend_model) + sibling_info = router.get_deployment_model_info(model_id=sibling_id, model_name=backend_model) + + assert custom_info is not None + assert custom_info["search_context_cost_per_query"]["search_context_size_low"] == 0.123 + assert litellm.model_cost[backend_model]["search_context_cost_per_query"] == expected_nested + assert sibling_info is not None + assert sibling_info["search_context_cost_per_query"] == expected_nested + finally: + _restore_model_cost_entries(model_keys) + litellm.get_model_info.cache_clear() + + +def test_base_model_custom_info_does_not_pollute_cached_base_model(): + base_model = "azure/gpt-4o" + deployment_id = "lit5471-base-model" + base_model_info = copy.deepcopy(litellm.get_model_info(model=base_model)) + model_keys = { + "azure/gpt-4o": copy.deepcopy(litellm.model_cost.get("azure/gpt-4o")), + deployment_id: copy.deepcopy(litellm.model_cost.get(deployment_id)), + } + try: + router = Router( + model_list=[ + { + "model_name": "azure-custom", + "litellm_params": { + "model": "gpt-4o", + "custom_llm_provider": "azure", + "api_key": "fake-key", + }, + "model_info": { + "id": deployment_id, + "base_model": base_model, + "input_cost_per_token": 0.777, + }, + } + ], + ) + + info = router.get_deployment_model_info(model_id=deployment_id, model_name=base_model) + + assert info is not None + assert info["input_cost_per_token"] == 0.777 + assert litellm.get_model_info(model=base_model) == base_model_info + finally: + _restore_model_cost_entries(model_keys) + litellm.get_model_info.cache_clear() + + +def test_builtin_only_deployment_info_is_not_the_cached_object(): + backend_model = "gpt-4o-search-preview" + deployment_id = "lit5471-builtin-only" + litellm.get_model_info.cache_clear() + model_keys = {deployment_id: copy.deepcopy(litellm.model_cost.get(deployment_id))} + try: + cached_info = litellm.get_model_info(model=backend_model) + assert cached_info["search_context_cost_per_query"] + + info = Router(model_list=[]).get_deployment_model_info(model_id=deployment_id, model_name=backend_model) + + assert info is not None + assert info["search_context_cost_per_query"] == cached_info["search_context_cost_per_query"] + assert _nested_container_ids(info).isdisjoint(_nested_container_ids(cached_info)) + finally: + _restore_model_cost_entries(model_keys) + litellm.get_model_info.cache_clear() + + +def test_custom_only_deployment_info_is_not_the_registry_entry(): + unknown_backend = "openai/lit5471-unknown-backend" + deployment_id = "lit5471-custom-only" + nested_pricing = {"search_context_size_low": 0.123} + model_keys = { + unknown_backend: copy.deepcopy(litellm.model_cost.get(unknown_backend)), + deployment_id: copy.deepcopy(litellm.model_cost.get(deployment_id)), + } + try: + router = Router( + model_list=[ + { + "model_name": "custom-only", + "litellm_params": {"model": unknown_backend, "api_key": "fake-key"}, + "model_info": {"id": deployment_id, "search_context_cost_per_query": dict(nested_pricing)}, + } + ], + ) + registry_entry = litellm.model_cost[deployment_id] + + info = router.get_deployment_model_info(model_id=deployment_id, model_name=unknown_backend) + + assert info is not None + assert info["search_context_cost_per_query"] == nested_pricing + assert _nested_container_ids(info).isdisjoint(_nested_container_ids(registry_entry)) + finally: + _restore_model_cost_entries(model_keys) + litellm.get_model_info.cache_clear() + + +def test_router_model_info_deep_copies_nested_cached_metadata(): + model = "openai/gpt-4o-search-preview" + litellm.get_model_info.cache_clear() + try: + cached_info = litellm.get_model_info(model=model) + assert cached_info is not None + expected_nested = copy.deepcopy(cached_info["search_context_cost_per_query"]) + assert expected_nested + + router = Router(model_list=[]) + merged_info = router.get_router_model_info( + deployment={ + "model_name": "search", + "litellm_params": {"model": "gpt-4o-search-preview"}, + "model_info": {"id": "lit5471-router-model-info"}, + }, + received_model_name="search", + ) + + assert merged_info["search_context_cost_per_query"] == expected_nested + assert _nested_container_ids(merged_info).isdisjoint(_nested_container_ids(cached_info)) + assert litellm.get_model_info(model=model)["search_context_cost_per_query"] == expected_nested + finally: + litellm.get_model_info.cache_clear() diff --git a/tests/test_litellm/test_test_quality_gate.py b/tests/test_litellm/test_test_quality_gate.py index 4caadca3d09..3bf4b89ac4e 100644 --- a/tests/test_litellm/test_test_quality_gate.py +++ b/tests/test_litellm/test_test_quality_gate.py @@ -1,9 +1,11 @@ """Tests for scripts/test_quality_gate.py. -The gate's whole value is that it blames a change only for what it adds, and that a -limit can never rise. Both properties live in pure functions, so they are tested -directly: `evaluate` for the blame rule, `ratcheted_budget` for the one-way ratchet, -and `parse_changed_lines` for the diff scan that turns a breach into file:line. +The gate's whole value is that it blames a change only for what it adds, that a limit +can never rise, and that a limit cannot stay above a count the branch pushed below it. +All three live in pure functions, so they are tested directly: `evaluate` for the blame +rule, `ratcheted_budget` for the one-way ratchet, `unratcheted` for the ceiling a branch +left behind, and `parse_changed_lines` for the diff scan that turns a breach into +file:line. """ import importlib.util @@ -66,11 +68,32 @@ def test_ratchet_never_goes_below_zero(): assert updated["TQ001"]["limit"] == 0 -def test_ratchet_leaves_a_rule_seeded_on_this_branch_untouched(): - updated = gate.ratcheted_budget( - _BUDGET, {"TQ001": 0}, {"TQ001": 10}, seeded=frozenset({"TQ001"}) - ) - assert updated["TQ001"]["limit"] == 10 +def test_ratchet_lowers_a_rule_introduced_on_this_branch_like_any_other(): + updated = gate.ratcheted_budget(_BUDGET, {"TQ001": 4}, {"TQ001": 10}) + assert updated["TQ001"]["limit"] == 4 + + +def test_a_branch_that_cleared_violations_must_lower_the_ceiling(): + stale = gate.unratcheted({"TQ001": 6}, {"TQ001": 10}, _BUDGET) + assert [(b.rule, b.total, b.cap, b.added) for b in stale] == [("TQ001", 6, 10, -4)] + + +def test_headroom_already_in_the_base_is_not_blamed_on_this_branch(): + assert gate.unratcheted({"TQ001": 6}, {"TQ001": 6}, _BUDGET) == () + + +def test_a_branch_that_cleared_down_to_the_ceiling_exactly_is_clean(): + assert gate.unratcheted({"TQ001": 10}, {"TQ001": 12}, _BUDGET) == () + + +def test_a_branch_that_added_violations_is_not_a_ratchet_finding(): + assert gate.unratcheted({"TQ001": 14}, {"TQ001": 10}, _BUDGET) == () + + +def test_the_ratchet_finding_survives_the_update_that_answers_it(): + cleared = {"TQ001": 6} + updated = gate.ratcheted_budget(_BUDGET, cleared, {"TQ001": 10}) + assert gate.unratcheted(cleared, {"TQ001": 10}, updated) == () def test_parse_changed_lines_groups_hunks_under_their_own_file(): @@ -121,5 +144,5 @@ def test_the_shipped_budget_covers_every_rule_the_checker_can_emit(): import json budget = json.loads((_REPO_ROOT / "test-quality-budget.json").read_text()) - assert set(budget) == {"TQ001", "TQ002", "TQ003", "TQ004", "TQ005", "TQ006"} + assert set(budget) == {"TQ001", "TQ002", "TQ003", "TQ004", "TQ005", "TQ006", "TQ007"} assert all(spec["limit"] >= 0 for spec in budget.values()) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index cb58038e081..041c60e0ba6 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -36,6 +36,7 @@ from litellm.utils import ( ProviderConfigManager, TextCompletionStreamWrapper, _check_provider_match, + _get_potential_model_names, _is_streaming_request, get_api_key, get_llm_provider, @@ -129,6 +130,74 @@ def test_get_model_info_surfaces_supports_adaptive_thinking(local_model_cost_map assert generalized["supports_adaptive_thinking"] is True +def test_potential_model_names_keeps_provider_prefixed_candidate(): + """A provider whose own model ids repeat the litellm provider name (Perplexity's + Agent API serves `perplexity/glm-5.2`, mapped as `perplexity/perplexity/glm-5.2`) + needs the un-stripped `/` candidate. Every other candidate reads + the leading `perplexity/` as the litellm prefix and strips it away.""" + already_prefixed = _get_potential_model_names( + model="perplexity/glm-5.2", custom_llm_provider="perplexity" + ) + assert already_prefixed["provider_prefixed_model_name"] == "perplexity/perplexity/glm-5.2" + assert already_prefixed["split_model"] == "glm-5.2" + assert already_prefixed["combined_model_name"] == "perplexity/glm-5.2" + assert already_prefixed["combined_stripped_model_name"] == "perplexity/glm-5.2" + + bare = _get_potential_model_names(model="glm-5.2", custom_llm_provider="perplexity") + assert bare["provider_prefixed_model_name"] == bare["combined_model_name"] == "perplexity/glm-5.2" + + +def test_get_model_info_resolves_provider_prefixed_model_ids(local_model_cost_map): + """Perplexity's Agent API third-party models are keyed `perplexity/perplexity/` + because Perplexity's own id already starts with `perplexity/`. Callers run + `get_llm_provider` first, which hands `_get_potential_model_names` model + `perplexity/glm-5.2` with provider `perplexity`, and every candidate but the + provider-prefixed one strips that second `perplexity/` off. Regression: the + entries were unreachable from `supports_reasoning` and from the cost calculator's + per-token fallback, so a mapped model reported no reasoning support and raised + "This model isn't mapped yet" on the only path where its rates are ever used.""" + for model, reasoning in ( + ("perplexity/perplexity/glm-5.2", True), + ("perplexity/perplexity/kimi-k3", True), + ("perplexity/perplexity/deepseek-v4-flash-0731", True), + ("perplexity/perplexity/kimi-k2.7-code", False), + ): + assert litellm.supports_reasoning(model=model) is reasoning, model + + via_provider = litellm.get_model_info( + model="perplexity/glm-5.2", custom_llm_provider="perplexity" + ) + assert via_provider["key"] == "perplexity/perplexity/glm-5.2" + assert via_provider["input_cost_per_token"] == 1.4e-06 + assert via_provider["output_cost_per_token"] == 4.4e-06 + assert via_provider["mode"] == "responses" + + +def test_provider_prefixed_lookup_never_outranks_an_existing_row(local_model_cost_map): + """The provider-prefixed candidate is tried last, after every candidate that + already existed, so no model that resolves today can change answer. `perplexity/sonar` + is the case that proves it: both `perplexity/sonar` and `perplexity/perplexity/sonar` + are cost-map keys, and the shorter one must keep winning.""" + sonar = litellm.get_model_info(model="sonar", custom_llm_provider="perplexity") + assert sonar["key"] == "perplexity/sonar" + assert sonar["mode"] == "chat" + assert sonar["input_cost_per_token"] == 1e-06 + + still_sonar = litellm.get_model_info( + model="perplexity/sonar", custom_llm_provider="perplexity" + ) + assert still_sonar["key"] == "perplexity/sonar" + assert still_sonar["mode"] == "chat" + + for model, provider, expected_key in ( + ("claude-sonnet-4-5", "anthropic", "claude-sonnet-4-5"), + ("anthropic/claude-sonnet-4-5", "anthropic", "claude-sonnet-4-5"), + ("gemini/gemini-2.0-flash", "gemini", "gemini/gemini-2.0-flash"), + ("openrouter/openai/gpt-4o", "openrouter", "openrouter/openai/gpt-4o"), + ): + assert litellm.get_model_info(model=model, custom_llm_provider=provider)["key"] == expected_key + + def test_check_provider_match_azure_ai_allows_openai_and_azure(): """ Test that azure_ai provider can match openai and azure models. @@ -3721,7 +3790,7 @@ class TestMetadataNoneHandling: # Attempting 'in' on None raises TypeError with pytest.raises(TypeError): - "model_group" in kwargs.get("metadata", {}) + _ = "model_group" in kwargs.get("metadata", {}) def test_litellm_params_metadata_none(self): """litellm_params.get("metadata") or {} should handle None value.""" diff --git a/tests/test_litellm/vector_stores/test_vector_store_create_provider_logic.py b/tests/test_litellm/vector_stores/test_vector_store_create_provider_logic.py index 5c279554c4e..bfb084e7dd2 100644 --- a/tests/test_litellm/vector_stores/test_vector_store_create_provider_logic.py +++ b/tests/test_litellm/vector_stores/test_vector_store_create_provider_logic.py @@ -36,7 +36,6 @@ def test_vector_store_create_with_simple_provider_name(): pytest.fail("Should not enter this branch for simple provider name") else: api_type = None - custom_llm_provider = custom_llm_provider # Keep as-is # Verify api_type is None assert api_type is None, "api_type should be None for simple provider names" @@ -132,7 +131,6 @@ def test_vector_store_create_with_ragflow_provider(): pytest.fail("Should not enter this branch for RAGFlow provider") else: api_type = None - custom_llm_provider = custom_llm_provider # Keep as-is # Verify api_type is None assert api_type is None, "api_type should be None for RAGFlow provider" diff --git a/tests/test_team.py b/tests/test_team.py index b1aba5c1311..62651beb6ec 100644 --- a/tests/test_team.py +++ b/tests/test_team.py @@ -511,10 +511,10 @@ async def test_team_update_sc_2(): print(f"team_data: {team_data}") ## assert rest of object is the same for k, v in new_team_data["data"].items(): - if ( - k == "members_with_roles" - ): # assert 1 more member (role: "user", user_email: $user_email) - len(new_team_data["data"][k]) == len(team_data[k]) + 1 + if k == "members_with_roles": + assert len(new_team_data["data"][k]) == len( + team_info["team_info"]["members_with_roles"] + ) elif ( k == "created_at" or k == "updated_at" diff --git a/type-discipline-budget.json b/type-discipline-budget.json index a5c5a9f135b..0e5d64fbb38 100644 --- a/type-discipline-budget.json +++ b/type-discipline-budget.json @@ -3,7 +3,7 @@ "limit": 22805 }, "LIT002": { - "limit": 26877 + "limit": 26873 }, "LIT003": { "limit": 269 @@ -27,7 +27,7 @@ "limit": 0 }, "LIT010": { - "limit": 16693 + "limit": 16673 }, "LIT011": { "limit": 5588 diff --git a/ui/litellm-dashboard/eslint-suppressions.json b/ui/litellm-dashboard/eslint-suppressions.json index 3bc93b7ebc4..b7c578d8ec6 100644 --- a/ui/litellm-dashboard/eslint-suppressions.json +++ b/ui/litellm-dashboard/eslint-suppressions.json @@ -1468,9 +1468,6 @@ }, "local/no-complex-jsx-arrow": { "count": 1 - }, - "react-hooks/set-state-in-effect": { - "count": 2 } }, "src/components/add_model/handle_add_auto_router_submit.tsx": { diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx index 7d94cae468d..8c36b934789 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.test.tsx @@ -2,6 +2,7 @@ import { fireEvent, render } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; import type { DailyData, KeyMetricWithMetadata, SpendMetrics } from "@/components/UsagePage/types"; +import type { DailyActivityRange } from "./useDailyActivityRange"; vi.mock("@/components/shared/advanced_date_picker", () => ({ __esModule: true, @@ -59,7 +60,7 @@ const dayWithModels = (date: string, models: Record +const renderWith = (results: DailyData[], overrides: Partial = {}) => render( results, loading: false, isFetchingMore: false, + progress: { currentPage: 1, totalPages: 1 }, + cancelled: false, + cancel: vi.fn(), + ...overrides, }} />, ); @@ -138,4 +143,38 @@ describe("CacheLeakageCard", () => { expect(getByText("No key usage in this range.")).toBeInTheDocument(); expect(queryByRole("table")).not.toBeInTheDocument(); }); + + it("tells the user the table is still filling in while fallback pages stream", () => { + const day = dayWithKeys("2026-07-12", { + "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), + }); + const { getByText, getByRole } = renderWith([day], { isFetchingMore: true }); + + expect(getByRole("table")).toBeInTheDocument(); + expect( + getByText("Data is still loading; rows and totals will update as the rest of the range arrives."), + ).toBeInTheDocument(); + }); + + it("keeps the streaming note off while a fresh range loads over the previous range's rows", () => { + const day = dayWithKeys("2026-07-12", { + "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), + }); + const { queryByText } = renderWith([day], { loading: true }); + + expect( + queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."), + ).not.toBeInTheDocument(); + }); + + it("drops the streaming note once the range has settled", () => { + const day = dayWithKeys("2026-07-12", { + "hash-leaky": key("leaky-key", { prompt_tokens: 10000, cache_read_input_tokens: 0 }), + }); + const { queryByText } = renderWith([day]); + + expect( + queryByText("Data is still loading; rows and totals will update as the rest of the range arrives."), + ).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx index ca47b71725d..3f27449ebe1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CacheLeakageCard.tsx @@ -123,6 +123,11 @@ const CacheLeakageCard: React.FC = ({ activity }) => { + {rows.length > 0 && isFetchingMore && ( +

+ Data is still loading; rows and totals will update as the rest of the range arrives. +

+ )} {rows.length === 0 ? (

{loading || isFetchingMore ? "Loading..." : `No ${emptyNoun} usage in this range.`} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx index 9d98233f110..2d46ca48adb 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.activity.test.tsx @@ -54,7 +54,7 @@ describe("CostOptimizationView daily activity", () => { useAuthorizedMock.mockReturnValue({ accessToken: "test-token", userId: "u1", userRole: "proxy_admin" }); const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); - const { getByRole, getByTestId, findByTestId } = render( + const { getByRole, getByTestId, findByTestId, queryByText } = render( , @@ -67,5 +67,28 @@ describe("CostOptimizationView daily activity", () => { expect(mockUserDailyActivityAggregatedCall).toHaveBeenCalledTimes(1); expect(mockUserDailyActivityCall).not.toHaveBeenCalled(); + expect(queryByText(/Currently fetching spend data/)).not.toBeInTheDocument(); + }); + + it("shows the fetch-progress banner while the paginated fallback streams pages in", async () => { + mockUserDailyActivityAggregatedCall.mockReset(); + mockUserDailyActivityCall.mockReset(); + mockUserDailyActivityAggregatedCall.mockRejectedValue(new Error("aggregated unavailable")); + mockUserDailyActivityCall.mockImplementation((...args: unknown[]) => + args[3] === 1 + ? Promise.resolve({ results: [], metadata: { total_pages: 3, has_more: true, page: 1 } }) + : new Promise(() => {}), + ); + useAuthorizedMock.mockReturnValue({ accessToken: "test-token", userId: "u1", userRole: "proxy_admin" }); + const queryClient = new QueryClient({ defaultOptions: { queries: { retry: false } } }); + + const { findByText, getByRole } = render( + + + , + ); + + expect(await findByText(/Currently fetching spend data: fetched 1 \/ 3 pages/)).toBeInTheDocument(); + expect(getByRole("button", { name: "Stop" })).toBeInTheDocument(); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx index 8122cfa7a9c..a588b28ecea 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx @@ -4,6 +4,7 @@ import React from "react"; import { Info, PiggyBank } from "lucide-react"; import useCan from "@/app/(dashboard)/hooks/useCan"; +import PaginationStatusAlerts from "@/components/shared/PaginationStatusAlerts"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import UsageTab from "./UsageTab"; import PromptCompressionTab from "./PromptCompressionTab"; @@ -62,6 +63,12 @@ const CostOptimizationView: React.FC = ({ accessToken

+ diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.test.tsx index a18109e8133..38517dab0ab 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.test.tsx @@ -33,6 +33,9 @@ describe("PromptCachingTab", () => { results: [], loading: false, isFetchingMore: false, + progress: { currentPage: 1, totalPages: 1 }, + cancelled: false, + cancel: vi.fn(), }; const { getByTestId } = render(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx index b72ebdc9c07..bceddf1eb7b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.test.tsx @@ -80,7 +80,9 @@ const job = (overrides: Partial = {}): ShadowEvalJob => ({ keys: [ { api_key_id: "hashed-key-abc", - max_turns: 200, + max_turns: 10000, + max_budget: 10, + spend: 3.21, stopped_at: null, key_alias: "prod-alpha", key_name: "sk-...alpha", @@ -133,7 +135,9 @@ const keyEntry = ( overrides: Partial = {}, ): ShadowEvalJob["keys"][number] => ({ api_key_id, - max_turns: 200, + max_turns: 10000, + max_budget: 10, + spend: 0, stopped_at: null, attempt_count: null, key_alias: null, @@ -321,6 +325,21 @@ describe("ShadowEvalSection", () => { expect(screen.getByText(/ends in 3 days/)).toBeInTheDocument(); }); + it("shows recorded eval spend against the job's dollar budget", () => { + const j = job(); + mockHooks({ jobs: [j], detailsById: { "job-1": j } }); + render(); + expect(screen.getByText(/\$3\.21 of \$10\.00 eval spend/)).toBeInTheDocument(); + }); + + it("shows spend without a budget cap for a job from before spend budgets existed", () => { + const j = job({ keys: [keyEntry("hashed-key-abc", { max_budget: null, spend: 3.21 })] }); + mockHooks({ jobs: [j], detailsById: { "job-1": j } }); + render(); + expect(screen.getByText(/\$3\.21 eval spend/)).toBeInTheDocument(); + expect(screen.queryByText(/of \$/)).not.toBeInTheDocument(); + }); + it("flags rows with fewer than 30 judged turns as low sample", () => { const j = job(); mockHooks({ jobs: [j], detailsById: { "job-1": j } }); @@ -388,7 +407,7 @@ describe("ShadowEvalSection", () => { direction: "forward", shadow_percentage: 10, duration_days: 7, - max_turns: 200, + max_budget: 10, judge_model: "anthropic/claude-sonnet-5", }; expect(start.mutate).toHaveBeenCalledWith(expectedBody); @@ -425,7 +444,7 @@ describe("ShadowEvalSection", () => { baseline_model: "prod-claude", shadow_percentage: 10, duration_days: 7, - max_turns: 200, + max_budget: 10, judge_model: "anthropic/claude-sonnet-5", }; expect(start.mutate).toHaveBeenCalledWith(expectedBody); @@ -463,8 +482,8 @@ describe("ShadowEvalSection", () => { job({ judged_count: 205, keys: [ - keyEntry("hash-spent", { max_turns: 200, stopped_at: "2026-08-08T00:00:00Z" }), - keyEntry("hash-hungry", { max_turns: 500 }), + keyEntry("hash-spent", { max_budget: 2, spend: 1.5, stopped_at: "2026-08-08T00:00:00Z" }), + keyEntry("hash-hungry", { max_budget: 5, spend: 0.2 }), ], results: { by_tier: [], @@ -492,25 +511,26 @@ describe("ShadowEvalSection", () => { if (!spent || !hungry) throw new Error("expected a table row per scoped key"); expect(within(spent).getByText("stopped")).toBeInTheDocument(); - expect(within(spent).getByText("200 / 200")).toBeInTheDocument(); + expect(within(spent).getByText("$1.50 / $2.00")).toBeInTheDocument(); expect(within(spent).getByText("60.0%")).toBeInTheDocument(); expect(within(hungry).getByText("running")).toBeInTheDocument(); - expect(within(hungry).getByText("0 / 500")).toBeInTheDocument(); + expect(within(hungry).getByText("$0.2000 / $5.00")).toBeInTheDocument(); expect(within(hungry).getByText("No verdicts yet")).toBeInTheDocument(); - expect(screen.getByText(/205 of 700 turns judged/)).toBeInTheDocument(); + expect(screen.getByText(/205 turns judged/)).toBeInTheDocument(); expect(screen.getByText(/Shadowing 10% of/)).toBeInTheDocument(); expect(screen.getByText("2 keys")).toBeInTheDocument(); }); it("reads a key that spent its budget as completed even before the sweep stamps it", () => { + const legacyTurnBudgetLeg = { max_budget: null, spend: 0.5, max_turns: 500, attempt_count: 3 }; mockHooks({ jobs: [ job({ keys: [ - keyEntry("hash-spent", { max_turns: 200, attempt_count: 200 }), - keyEntry("hash-hungry", { max_turns: 500, attempt_count: 3 }), + keyEntry("hash-spent", { max_budget: 2, spend: 2, attempt_count: 40 }), + keyEntry("hash-hungry", legacyTurnBudgetLeg), ], }), ], @@ -521,9 +541,9 @@ describe("ShadowEvalSection", () => { const hungry = screen.getByText("hash-hungr…").closest("tr"); if (!spent || !hungry) throw new Error("expected a table row per scoped key"); expect(within(spent).getByText("completed")).toBeInTheDocument(); - expect(within(spent).getByText("200 / 200")).toBeInTheDocument(); + expect(within(spent).getByText("$2.00 / $2.00")).toBeInTheDocument(); expect(within(hungry).getByText("running")).toBeInTheDocument(); - expect(within(hungry).getByText("3 / 500")).toBeInTheDocument(); + expect(within(hungry).getByText("3 / 500 turns")).toBeInTheDocument(); }); it("shows the per-key table while a multi-key job is still collecting, before any verdicts exist", () => { @@ -533,8 +553,8 @@ describe("ShadowEvalSection", () => { judged_count: 0, results: null, keys: [ - keyEntry("hash-spent", { max_turns: 2, attempt_count: 2 }), - keyEntry("hash-hungry", { max_turns: 500, attempt_count: 1 }), + keyEntry("hash-spent", { max_budget: 0.5, spend: 0.5, attempt_count: 2 }), + keyEntry("hash-hungry", { max_budget: 5, spend: 0.01, attempt_count: 1 }), ], }), ], @@ -544,7 +564,7 @@ describe("ShadowEvalSection", () => { const spent = screen.getByText("hash-spent…").closest("tr"); if (!spent) throw new Error("expected a per-key row before verdicts exist"); expect(within(spent).getByText("completed")).toBeInTheDocument(); - expect(within(spent).getByText("2 / 2")).toBeInTheDocument(); + expect(within(spent).getByText("$0.5000 / $0.5000")).toBeInTheDocument(); expect(screen.getByText("Budget used")).toBeInTheDocument(); expect(screen.queryByText("Judged turns")).not.toBeInTheDocument(); expect(screen.getByText(/Collecting verdicts/)).toBeInTheDocument(); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx index bc5044feaa6..44054e3b7c4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/ShadowEvalSection.tsx @@ -57,9 +57,19 @@ export const shadowedKeyLabel = (key: ShadowEvalJobKey): string => const shadowedKeysLabel = (job: ShadowEvalJob): string => job.keys.length === 1 ? shadowedKeyLabel(job.keys[0]) : `${job.keys.length} keys`; -const totalBudget = (job: ShadowEvalJob): number => job.keys.reduce((sum, key) => sum + key.max_turns, 0); +const totalBudget = (job: ShadowEvalJob): number | null => + job.keys.reduce( + (sum, key) => (sum === null || key.max_budget == null ? null : sum + key.max_budget), + 0, + ); -const keySpent = (key: ShadowEvalJobKey): boolean => key.attempt_count != null && key.attempt_count >= key.max_turns; +const totalSpend = (job: ShadowEvalJob): number => job.keys.reduce((sum, key) => sum + (key.spend ?? 0), 0); + +const keySpent = (key: ShadowEvalJobKey): boolean => { + const spendBudgetReached = key.max_budget != null && key.spend != null && key.spend >= key.max_budget; + const turnValveReached = key.attempt_count != null && key.attempt_count >= key.max_turns; + return spendBudgetReached || turnValveReached; +}; const keyStatus = (job: ShadowEvalJob, key: ShadowEvalJobKey): string => { if (job.status === "completed" || (key.stopped_at == null && keySpent(key))) return "completed"; @@ -207,7 +217,9 @@ const KeyTable: React.FC<{ job: ShadowEvalJob }> = ({ job }) => { - {(key.attempt_count ?? slice?.turn_count ?? 0).toLocaleString()} / {key.max_turns.toLocaleString()} + {key.max_budget != null + ? `${usd(key.spend ?? 0)} / ${usd(key.max_budget)}` + : `${(key.attempt_count ?? slice?.turn_count ?? 0).toLocaleString()} / ${key.max_turns.toLocaleString()} turns`} {slice ? ( <> @@ -300,8 +312,9 @@ const JobResults: React.FC<{

{jobHeadline(job)}

- {(job.judged_count ?? 0).toLocaleString()} of {totalBudget(job).toLocaleString()} turns judged ·{" "} - {(job.error_count ?? 0).toLocaleString()} errored · {usd(job.judge_spend ?? 0)} judge spend + {(job.judged_count ?? 0).toLocaleString()} turns judged · {(job.error_count ?? 0).toLocaleString()}{" "} + errored · {usd(totalSpend(job))} + {totalBudget(job) !== null ? ` of ${usd(totalBudget(job) ?? 0)}` : ""} eval spend {active && remaining ? ` · ${remaining}` : ""}

@@ -375,9 +388,9 @@ const DIRECTION_OPTIONS: readonly { value: ShadowEvalDirection; label: string }[ const START_FORM_DESCRIPTION: Record = { forward: - "Duplicates a sampled slice of the selected keys' traffic through the auto-router and has an LLM judge compare both answers blind. Each key gets its own turn budget. The router's answers are never served to users; judge calls bill to the shadowed key.", + "Duplicates a sampled slice of the selected keys' traffic through the auto-router and has an LLM judge compare both answers blind. Each key gets its own spend budget. The router's answers are never served to users; judge calls bill to the shadowed key.", reverse: - "Duplicates a sampled slice of the traffic the auto-router already serves against a fixed baseline model and has an LLM judge compare both answers blind. Each key gets its own turn budget. The baseline's answers are never served to users; judge calls bill to the shadowed key.", + "Duplicates a sampled slice of the traffic the auto-router already serves against a fixed baseline model and has an LLM judge compare both answers blind. Each key gets its own spend budget. The baseline's answers are never served to users; judge calls bill to the shadowed key.", }; const DURATION_OPTIONS = [ @@ -445,7 +458,7 @@ const StartForm: React.FC = () => { const [percentage, setPercentage] = useState("10"); const [durationDays, setDurationDays] = useState("7"); const [judgeModel, setJudgeModel] = useState(""); - const [maxTurns, setMaxTurns] = useState("200"); + const [maxBudget, setMaxBudget] = useState("10"); const { data: autoRouters } = useAutoRouters(); const judgeModelOptions = useJudgeModelOptions(); const baselineModelOptions = useBaselineModelOptions(); @@ -460,11 +473,11 @@ const StartForm: React.FC = () => { const parsedPct = Number.parseFloat(percentage); const percentageValid = parsedPct >= 0.1 && parsedPct <= 100; - const parsedMaxTurns = Number.parseInt(maxTurns, 10); - const maxTurnsValid = parsedMaxTurns >= 1 && parsedMaxTurns <= 2000; + const parsedMaxBudget = Number.parseFloat(maxBudget); + const maxBudgetValid = parsedMaxBudget >= 0.01 && parsedMaxBudget <= 10000; const baselinePicked = direction === "forward" || baselineModel !== ""; const filled = apiKeyIds.length > 0 && [routerName, judgeModel].every((field) => field !== "") && baselinePicked; - const boundsValid = percentageValid && maxTurnsValid; + const boundsValid = percentageValid && maxBudgetValid; const valid = Boolean(accessToken) && filled && boundsValid; const handleStart = () => { const startBody = { @@ -474,7 +487,7 @@ const StartForm: React.FC = () => { ...(direction === "reverse" ? { baseline_model: baselineModel } : {}), shadow_percentage: parsedPct, duration_days: Number.parseInt(durationDays, 10), - max_turns: parsedMaxTurns, + max_budget: parsedMaxBudget, judge_model: judgeModel, }; start.mutate(startBody); @@ -551,20 +564,22 @@ const StartForm: React.FC = () => { - +
+ $ setMaxTurns(e.target.value)} + value={maxBudget} + onChange={(e) => setMaxBudget(e.target.value)} /> - turns judged, max + max shadow + judge spend, per key
- {maxTurns.trim() !== "" && !maxTurnsValid && ( -

Enter a value from 1 to 2000

+ {maxBudget.trim() !== "" && !maxBudgetValid && ( +

Enter a value from 0.01 to 10000

)}
{direction === "reverse" && ( @@ -620,7 +635,7 @@ const PreviousJob: React.FC<{ job: ShadowEvalJob }> = ({ job }) => {

{jobHeadline(shown)}

{shown.judged_count != null && - `${shown.judged_count.toLocaleString()} judged · ${(shown.error_count ?? 0).toLocaleString()} errored · ${usd(shown.judge_spend ?? 0)} judge spend · `} + `${shown.judged_count.toLocaleString()} judged · ${(shown.error_count ?? 0).toLocaleString()} errored · ${usd(totalSpend(shown))} eval spend · `} {new Date(shown.created_at).toLocaleDateString()}

diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx index 20d57754857..74a936369c9 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx @@ -121,6 +121,9 @@ const renderWith = (results: DailyData[], options: RenderOptions = {}) => { results, loading: false, isFetchingMore: false, + progress: { currentPage: 1, totalPages: 1 }, + cancelled: false, + cancel: vi.fn(), }} />, ); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx index 43c4aa04e2b..2229438d844 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.test.tsx @@ -3,10 +3,19 @@ import { describe, expect, it, vi } from "vitest"; const mockUsePaginatedDailyActivity = vi.fn(); +const mockCancel = vi.fn(); + vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", () => ({ usePaginatedDailyActivity: (args: unknown) => { mockUsePaginatedDailyActivity(args); - return { data: { results: [] }, loading: false, isFetchingMore: false }; + return { + data: { results: [] }, + loading: false, + isFetchingMore: false, + progress: { currentPage: 4, totalPages: 9 }, + cancelled: false, + cancel: mockCancel, + }; }, })); @@ -41,6 +50,14 @@ describe("useDailyActivityRange", () => { ); }); + it("forwards the pagination progress and cancel affordances instead of dropping them", () => { + const { result } = renderHook(() => useDailyActivityRange("test-token", "u1", "proxy_admin")); + + expect(result.current.progress).toEqual({ currentPage: 4, totalPages: 9 }); + expect(result.current.cancelled).toBe(false); + expect(result.current.cancel).toBe(mockCancel); + }); + it("stays disabled until an access token is available", () => { renderHook(() => useDailyActivityRange(null, "u1", "proxy_admin")); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts index e16458728a1..81ddb6af585 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/useDailyActivityRange.ts @@ -18,6 +18,9 @@ export interface DailyActivityRange { results: DailyData[]; loading: boolean; isFetchingMore: boolean; + progress: { currentPage: number; totalPages: number }; + cancelled: boolean; + cancel: () => void; } export const useDailyActivityRange = ( @@ -33,7 +36,7 @@ export const useDailyActivityRange = ( const endTime = dateValue.to ?? null; const effectiveUserId = all_admin_roles.includes(userRole) ? null : userId; - const { data, loading, isFetchingMore } = usePaginatedDailyActivity({ + const { data, loading, isFetchingMore, progress, cancelled, cancel } = usePaginatedDailyActivity({ fetchFn: userDailyActivityCall, aggregatedFetchFn: userDailyActivityAggregatedCall, args: [accessToken, startTime, endTime, effectiveUserId, true], @@ -46,5 +49,8 @@ export const useDailyActivityRange = ( results: data.results as DailyData[], loading, isFetchingMore, + progress, + cancelled, + cancel, }; }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx index 69198e42279..9501fa7a9a1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/usage/_components/components/EntityUsage/EntityUsage.tsx @@ -15,10 +15,9 @@ import { Card as ShadcnCard, CardContent, CardHeader, CardTitle } from "@/compon import { hasCapability, type Capability } from "@/utils/capabilities"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import type { DateRangePickerValue } from "@/components/shared/date_picker_types"; -import { ChevronDown, ChevronRight, ExternalLink, Info, Loader2 } from "lucide-react"; +import { ChevronDown, ChevronRight, Info } from "lucide-react"; import type { ColumnDef } from "@tanstack/react-table"; -import { Alert, AlertDescription } from "@/components/shared/Alert"; -import { Button } from "@/components/ui/button"; +import PaginationStatusAlerts from "@/components/shared/PaginationStatusAlerts"; import { Tabs, TabsContent, TabsList, TabsTrigger } from "@/components/ui/tabs"; import { Tooltip, TooltipContent, TooltipTrigger } from "@/components/ui/tooltip"; import React, { type ReactNode, useMemo, useState } from "react"; @@ -643,57 +642,20 @@ const EntityUsage: React.FC = ({ return (
- {isFetchingMore && ( - - - - - Currently fetching spend data: fetched {progress.currentPage} / {progress.totalPages} pages. Charts will - update periodically as data loads. Moving off of this page will stop and reset this. To continue using the - UI in the meantime,{" "} - - open a new tab - - . - - - - - )} - {cancelled && ( - - - Showing partial data ({progress.currentPage}/{progress.totalPages} pages loaded) - - - )} - {agentIsFetchingMore && showAgentBreakdown && ( - - - - - Currently fetching agent data: fetched {agentProgress.currentPage} / {agentProgress.totalPages} pages. - Charts will update periodically as data loads. Moving off of this page will stop and reset this. To - continue using the UI in the meantime,{" "} - - open a new tab - - . - - - - - )} - {agentCancelled && showAgentBreakdown && ( - - - Showing partial agent data ({agentProgress.currentPage}/{agentProgress.totalPages} pages loaded) - - + + {showAgentBreakdown && ( + )} = ({ teams, organizations }) => { />
- {paginatedResult.isFetchingMore && ( - - - - - Currently fetching spend data: fetched {paginatedResult.progress.currentPage} /{" "} - {paginatedResult.progress.totalPages} pages. Charts will update periodically as data loads. Moving off - of this page will stop and reset this. To continue using the UI in the meantime,{" "} - - open a new tab - - . - - - - - )} - {paginatedResult.cancelled && ( - - - Showing partial data ({paginatedResult.progress.currentPage}/{paginatedResult.progress.totalPages} pages - loaded) - - - )} + {/* Your Usage / Global Usage Panel */} {(usageView === "global" || usageView === "my-usage") && ( <> diff --git a/ui/litellm-dashboard/src/components/ThemeToggle/ThemeToggle.test.tsx b/ui/litellm-dashboard/src/components/ThemeToggle/ThemeToggle.test.tsx index 3f7af94bfa0..41efedcd43a 100644 --- a/ui/litellm-dashboard/src/components/ThemeToggle/ThemeToggle.test.tsx +++ b/ui/litellm-dashboard/src/components/ThemeToggle/ThemeToggle.test.tsx @@ -16,7 +16,7 @@ const openMenu = async () => { await screen.findByRole("menu"); }; -const pick = async (label: string) => userEvent.click(screen.getByRole("menuitemradio", { name: label })); +const pick = async (label: string | RegExp) => userEvent.click(screen.getByRole("menuitemradio", { name: label })); beforeEach(() => { localStorage.clear(); @@ -33,7 +33,7 @@ describe("ThemeToggle", () => { await openMenu(); expect(screen.getByRole("menuitemradio", { name: "Light" })).toBeChecked(); - expect(screen.getByRole("menuitemradio", { name: "Dark" })).not.toBeChecked(); + expect(screen.getByRole("menuitemradio", { name: /^Dark/ })).not.toBeChecked(); expect(screen.getByRole("menuitemradio", { name: "System" })).not.toBeChecked(); }); @@ -41,7 +41,7 @@ describe("ThemeToggle", () => { renderToggle(); await openMenu(); - await pick("Dark"); + await pick(/^Dark/); expect(document.documentElement).toHaveClass("dark"); expect(localStorage.getItem("theme")).toBe("dark"); @@ -50,7 +50,7 @@ describe("ThemeToggle", () => { it("hands control back to the system preference when asked", async () => { renderToggle(); await openMenu(); - await pick("Dark"); + await pick(/^Dark/); await pick("System"); @@ -58,15 +58,20 @@ describe("ThemeToggle", () => { expect(document.documentElement).not.toHaveClass("dark"); }); - it("flags dark mode as experimental, and only while it is on", async () => { + it("marks dark as beta in the menu, and leaves the other choices unmarked", async () => { renderToggle(); await openMenu(); - expect(screen.queryByText("Experimental")).not.toBeInTheDocument(); - await pick("Dark"); - expect(screen.getByText("Experimental")).toBeInTheDocument(); + expect(screen.getByRole("menuitemradio", { name: /^Dark/ })).toHaveTextContent("Beta"); + expect(screen.getByRole("menuitemradio", { name: "Light" })).not.toHaveTextContent("Beta"); + expect(screen.getByRole("menuitemradio", { name: "System" })).not.toHaveTextContent("Beta"); + }); - await pick("Light"); - expect(screen.queryByText("Experimental")).not.toBeInTheDocument(); + it("keeps the beta marker inside the menu rather than in the toolbar", async () => { + renderToggle(); + await openMenu(); + await pick(/^Dark/); + + expect(screen.getByRole("button", { name: "Theme" })).not.toHaveTextContent("Beta"); }); }); diff --git a/ui/litellm-dashboard/src/components/ThemeToggle/ThemeToggle.tsx b/ui/litellm-dashboard/src/components/ThemeToggle/ThemeToggle.tsx index 6a006141523..3fbb3d2eb5b 100644 --- a/ui/litellm-dashboard/src/components/ThemeToggle/ThemeToggle.tsx +++ b/ui/litellm-dashboard/src/components/ThemeToggle/ThemeToggle.tsx @@ -15,46 +15,43 @@ import { } from "@/components/ui/dropdown-menu"; const THEMES = [ - { value: "system", label: "System", Icon: Monitor }, - { value: "light", label: "Light", Icon: Sun }, - { value: "dark", label: "Dark", Icon: Moon }, + { value: "system", label: "System", Icon: Monitor, beta: false }, + { value: "light", label: "Light", Icon: Sun, beta: false }, + { value: "dark", label: "Dark", Icon: Moon, beta: true }, ] as const; const ThemeToggle: React.FC = () => { const { theme, setTheme, resolvedTheme } = useTheme(); - const isDark = resolvedTheme === "dark"; return ( - - {isDark && ( - - Experimental - - )} - - - } - > - {isDark ? : } - - - - {THEMES.map(({ value, label, Icon }) => ( - - - {label} - - ))} - - - - + + + } + > + {resolvedTheme === "dark" ? : } + + + + {THEMES.map(({ value, label, Icon, beta }) => ( + + + {label} + {beta && ( + + Beta + + )} + + ))} + + + ); }; diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx index 640b10ad163..0a848ed7deb 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.test.tsx @@ -8,9 +8,9 @@ vi.mock( ); const mockModelInfo = [ - { model_group: "gpt-4", mode: "chat" }, + { model_group: "gpt-4", mode: "chat", supports_reasoning: true }, { model_group: "gpt-3.5-turbo", mode: "chat" }, - { model_group: "claude-3-opus", mode: "chat" }, + { model_group: "claude-3-opus", mode: "chat", supports_reasoning: true }, { model_group: "text-embedding-3-small", mode: "embedding" }, ] as any[]; @@ -403,7 +403,7 @@ describe("ComplexityRouterConfig", () => { renderWithProviders(); const simpleTierSection = screen.getByText("Simple Tier").closest(".mb-4") as HTMLElement; - const combobox = within(simpleTierSection).getByRole("combobox"); + const combobox = within(simpleTierSection).getByRole("combobox", { name: "Select model(s) for simple queries" }); await user.click(combobox); expect((await screen.findAllByText("gpt-3.5-turbo")).length).toBeGreaterThan(0); @@ -874,3 +874,72 @@ describe("plan-mode override", () => { expect(await screen.findByRole("switch", { name: switchName })).toHaveAttribute("aria-disabled", "true"); }); }); + +describe("ComplexityRouterConfig per-model reasoning effort", () => { + it("renders one effort select per selected model, defaulting to Default", () => { + renderWithProviders(); + const select = screen.getByRole("combobox", { name: "Reasoning effort for gpt-4 in the Complex tier" }); + expect(select).toHaveTextContent("Default"); + }); + + it("shows the hydrated effort for a model that has one stored", () => { + renderWithProviders( + , + ); + const select = screen.getByRole("combobox", { name: "Reasoning effort for gpt-4 in the Complex tier" }); + expect(select).toHaveTextContent("high"); + }); + + it("emits tier_model_params scoped to the tier and model when an effort is picked", async () => { + const onChange = vi.fn(); + renderWithProviders(); + const user = userEvent.setup(); + await user.click(screen.getByRole("combobox", { name: "Reasoning effort for gpt-4 in the Complex tier" })); + await user.click(await screen.findByRole("option", { name: "high" })); + expect(onChange).toHaveBeenCalledWith({ + ...defaultValue, + tier_model_params: { COMPLEX: { "gpt-4": { reasoning_effort: "high" } } }, + }); + }); + + it("picking Default removes the stored effort", async () => { + const onChange = vi.fn(); + renderWithProviders( + , + ); + const user = userEvent.setup(); + await user.click(screen.getByRole("combobox", { name: "Reasoning effort for gpt-4 in the Complex tier" })); + await user.click(await screen.findByRole("option", { name: "Default" })); + expect(onChange).toHaveBeenCalledWith({ ...defaultValue, tier_model_params: undefined }); + }); +}); + +describe("ComplexityRouterConfig reasoning effort gating", () => { + it("offers no effort select for a model group without reasoning support", () => { + renderWithProviders(); + expect( + screen.queryByRole("combobox", { name: "Reasoning effort for gpt-3.5-turbo in the Simple tier" }), + ).not.toBeInTheDocument(); + }); + + // A stored effort on a model the group info calls non-reasoning must stay visible, or the + // operator has no way to clear it. + it("keeps the select for a non-reasoning model that already has a stored effort", () => { + renderWithProviders( + , + ); + expect( + screen.getByRole("combobox", { name: "Reasoning effort for gpt-3.5-turbo in the Simple tier" }), + ).toHaveTextContent("low"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx index 57028273d6d..fc731e2c77f 100644 --- a/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx +++ b/ui/litellm-dashboard/src/components/add_model/ComplexityRouterConfig.tsx @@ -12,7 +12,15 @@ import React from "react"; import { ModelGroup } from "@/components/llm_calls/fetch_models"; import AdaptiveRoutingConfig from "./AdaptiveRoutingConfig"; import ClassificationMethodConfig from "./ClassificationMethodConfig"; -import { resolveComplexityDefaultModel, tierOptions } from "./complexity_router_tiers"; +import { + ReasoningEffort, + TierModelParamsByTier, + pruneTierModelParams, + resolveComplexityDefaultModel, + setTierModelReasoningEffort, + tierOptions, +} from "./complexity_router_tiers"; +import TierModelEffortRows from "./TierModelEffortRows"; import EscalationKeywords from "./EscalationKeywords"; import KeywordTierRules, { KeywordTierRule } from "./KeywordTierRules"; import SemanticKeywordMatching from "./SemanticKeywordMatching"; @@ -153,6 +161,12 @@ export interface ComplexityRouterConfigValue { * floor tracks tier_boundaries.simple_medium; an explicit 0 is a real floor that promotes on the markers alone. */ reasoning_override_min_score?: number; + /** + * Per-(tier, model) litellm_params, serialized to the sibling tier_model_configs key. The full + * params object is held, not just reasoning_effort, so keys authored in config.yaml survive an + * edit round-trip. + */ + tier_model_params?: TierModelParamsByTier; } interface ComplexityRouterConfigProps { @@ -237,6 +251,10 @@ const ComplexityRouterConfig: React.FC = ({ const defaultModel = resolveComplexityDefaultModel(value.tiers, value.default_model); // Embedding models can't serve a chat-completion role, so they're excluded here. + const reasoningModels = new Set( + modelInfo.filter((model) => model.supports_reasoning).map((model) => model.model_group), + ); + const modelOptions = modelInfo .filter((model) => model.mode !== "embedding") .map((model) => ({ @@ -248,6 +266,18 @@ const ComplexityRouterConfig: React.FC = ({ onChange({ ...value, tiers: { ...value.tiers, [tier]: models }, + tier_model_params: pruneTierModelParams(value.tier_model_params, tier, models), + }); + }; + + const handleTierModelEffortChange = ( + tier: keyof ComplexityTiers, + model: string, + effort: ReasoningEffort | undefined, + ) => { + onChange({ + ...value, + tier_model_params: setTierModelReasoningEffort(value.tier_model_params, tier, model, effort), }); }; @@ -332,6 +362,13 @@ const ComplexityRouterConfig: React.FC = ({ emptyText="No models found" className={tierMissing ? "w-full border-destructive" : "w-full"} /> + handleTierModelEffortChange(tier, model, effort)} + /> {value.tiers[tier].length > 1 && ( Multiple models selected — the router randomly picks among them per request (or Thompson-samples diff --git a/ui/litellm-dashboard/src/components/add_model/TierModelEffortRows.tsx b/ui/litellm-dashboard/src/components/add_model/TierModelEffortRows.tsx new file mode 100644 index 00000000000..67f583bd894 --- /dev/null +++ b/ui/litellm-dashboard/src/components/add_model/TierModelEffortRows.tsx @@ -0,0 +1,80 @@ +import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from "@/components/ui/select"; +import { SimpleTooltip } from "@/components/ui/tooltip"; +import { Info } from "lucide-react"; +import React from "react"; +import { REASONING_EFFORT_OPTIONS, ReasoningEffort, TierModelParams } from "./complexity_router_tiers"; + +const PROVIDER_DEFAULT = "__provider_default__"; + +const asEffort = (params: TierModelParams | undefined): ReasoningEffort | undefined => { + const stored = params?.reasoning_effort; + if (typeof stored !== "string") return undefined; + return REASONING_EFFORT_OPTIONS.find((option) => option === stored); +}; + +interface TierModelEffortRowsProps { + tierLabel: string; + models: string[]; + reasoningModels: ReadonlySet; + paramsByModel: Record | undefined; + onEffortChange: (model: string, effort: ReasoningEffort | undefined) => void; +} + +const TierModelEffortRows: React.FC = ({ + tierLabel, + models, + reasoningModels, + paramsByModel, + onEffortChange, +}) => { + const shown = models.filter( + (model) => reasoningModels.has(model) || Object.keys(paramsByModel?.[model] ?? {}).length > 0, + ); + if (shown.length === 0) return null; + return ( +
+
+ Reasoning effort + + + +
+ {shown.map((model) => ( +
+ {model} + +
+ ))} +
+ ); +}; + +export default TierModelEffortRows; diff --git a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx index 71fd6f4f1f9..46ab0bc2c9f 100644 --- a/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx +++ b/ui/litellm-dashboard/src/components/add_model/add_auto_router_tab.tsx @@ -368,6 +368,7 @@ const AddAutoRouterTab: React.FC = ({ tierDistancePenalty: complexityRouterConfig.tier_distance_penalty ?? DEFAULT_TIER_DISTANCE_PENALTY, adaptiveEligible: complexityRouterConfig.adaptive_eligible ?? "all", returnRawModelName: complexityRouterConfig.return_raw_model_name ?? false, + tierModelParams: complexityRouterConfig.tier_model_params, tierBoundaries: complexityRouterConfig.tier_boundaries, tokenThresholds: complexityRouterConfig.token_thresholds, dimensionWeights: complexityRouterConfig.dimension_weights, diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts index 3e33136a3cc..cb55362c6a7 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.test.ts @@ -656,3 +656,20 @@ describe("getPlanModeTierError", () => { expect(getPlanModeTierError("COMPLEX", tiersWithEmptyComplex)).toContain("COMPLEX"); }); }); + +describe("buildComplexityRouterConfig tier model params", () => { + it("keeps tier_model_configs out of the payload when nothing is set", () => { + expect(buildComplexityRouterConfig(baseParams)).not.toHaveProperty("tier_model_configs"); + }); + + it("emits tier_model_configs beside string tiers when efforts are set", () => { + const config = buildComplexityRouterConfig({ + ...baseParams, + tierModelParams: { COMPLEX: { "claude-sonnet-4": { reasoning_effort: "high" } } }, + }); + expect(config.tiers).toEqual(tiers); + expect(config.tier_model_configs).toEqual({ + COMPLEX: [{ model_name: "claude-sonnet-4", litellm_params: { reasoning_effort: "high" } }], + }); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts index bddf8321ad2..677cbe7063f 100644 --- a/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts +++ b/ui/litellm-dashboard/src/components/add_model/build_complexity_router_config.ts @@ -1,5 +1,6 @@ import { KeywordTierRule } from "./KeywordTierRules"; import { emptyKeywordTierRuleIndexes, serializeKeywordTierRules } from "./complexity_router_keywords"; +import { TierModelParams, TierModelParamsByTier, serializeTierModelConfigs } from "./complexity_router_tiers"; import { AdaptiveEligible, AdaptiveRouterWeights, @@ -99,6 +100,7 @@ export interface BuildComplexityRouterConfigParams { tokenThresholds?: TokenThresholds; dimensionWeights?: DimensionWeights; reasoningOverrideMinScore?: number; + tierModelParams?: TierModelParamsByTier; } export interface ComplexityRouterConfigPayload { @@ -129,6 +131,7 @@ export interface ComplexityRouterConfigPayload { token_thresholds?: TokenThresholds; dimension_weights?: DimensionWeights; reasoning_override_min_score?: number; + tier_model_configs?: Record; } const TIER_KEYS: Array = ["SIMPLE", "MEDIUM", "COMPLEX", "REASONING"]; @@ -236,7 +239,9 @@ export const buildComplexityRouterConfig = ({ tokenThresholds, dimensionWeights, reasoningOverrideMinScore, + tierModelParams, }: BuildComplexityRouterConfigParams): ComplexityRouterConfigPayload => { + const serializedTierModelConfigs = serializeTierModelConfigs(tiers, tierModelParams); const cleanedEscalationKeywords = escalationKeywords.map((keyword) => keyword.trim()).filter(Boolean); const cleanedKeywordTierRules = serializeKeywordTierRules(keywordTierRules); const cleanedTierLabels = serializeTierLabels(tierLabels); @@ -252,6 +257,7 @@ export const buildComplexityRouterConfig = ({ return { tiers, + ...(serializedTierModelConfigs && { tier_model_configs: serializedTierModelConfigs }), ...(defaultModel?.trim() && { default_model: defaultModel }), ...(planModeMinTier?.trim() && { plan_mode_min_tier: planModeMinTier }), ...(cleanedTierLabels && { tier_labels: cleanedTierLabels }), diff --git a/ui/litellm-dashboard/src/components/add_model/complexity_router_tiers.test.ts b/ui/litellm-dashboard/src/components/add_model/complexity_router_tiers.test.ts index 8f75b26f250..4dffbbd2ac7 100644 --- a/ui/litellm-dashboard/src/components/add_model/complexity_router_tiers.test.ts +++ b/ui/litellm-dashboard/src/components/add_model/complexity_router_tiers.test.ts @@ -1,6 +1,13 @@ import { describe, expect, it } from "vitest"; -import { normalizeTierModels, resolveComplexityDefaultModel } from "./complexity_router_tiers"; +import { + hydrateTierModelParams, + normalizeTierModels, + pruneTierModelParams, + resolveComplexityDefaultModel, + serializeTierModelConfigs, + setTierModelReasoningEffort, +} from "./complexity_router_tiers"; import type { ComplexityTiers } from "./ComplexityRouterConfig"; @@ -70,3 +77,160 @@ describe("resolveComplexityDefaultModel", () => { expect(resolveComplexityDefaultModel(noTiers)).toBeUndefined(); }); }); + +// The backend also accepts `{model_name, litellm_params}` entries and splits them into the +// sibling tier_model_configs key at validation (config.py `_normalize_tier_model_configs`). +// Before this widening, an object entry was silently dropped here, so opening the edit modal on +// a yaml-authored config rendered the tier empty and the next save destroyed it. +describe("normalizeTierModels object entries", () => { + it("reads model_name from an object entry the way the backend does", () => { + expect(normalizeTierModels([{ model_name: "opus", litellm_params: { reasoning_effort: "high" } }, "mini"])).toEqual( + ["opus", "mini"], + ); + }); + + it("widens a single object entry to a one-element pool", () => { + expect(normalizeTierModels({ model_name: "opus" })).toEqual(["opus"]); + }); + + it("drops an object without a model_name", () => { + expect(normalizeTierModels([{ litellm_params: { reasoning_effort: "high" } }])).toEqual([]); + }); +}); + +describe("hydrateTierModelParams", () => { + it("reads the sibling tier_model_configs key", () => { + expect( + hydrateTierModelParams( + { MEDIUM: ["opus"] }, + { MEDIUM: [{ model_name: "opus", litellm_params: { reasoning_effort: "medium" } }] }, + ), + ).toEqual({ MEDIUM: { opus: { reasoning_effort: "medium" } } }); + }); + + it("reads inline object entries out of tiers", () => { + expect( + hydrateTierModelParams( + { COMPLEX: [{ model_name: "opus", litellm_params: { reasoning_effort: "high" } }] }, + undefined, + ), + ).toEqual({ COMPLEX: { opus: { reasoning_effort: "high" } } }); + }); + + // config.py merges the two sources with tier_model_configs winning per (tier, model); hydrating + // the other way round would show the operator a value the router never uses. + it("lets tier_model_configs beat an inline entry for the same tier and model", () => { + expect( + hydrateTierModelParams( + { MEDIUM: [{ model_name: "opus", litellm_params: { reasoning_effort: "low" } }] }, + { MEDIUM: [{ model_name: "opus", litellm_params: { reasoning_effort: "medium" } }] }, + ), + ).toEqual({ MEDIUM: { opus: { reasoning_effort: "medium" } } }); + }); + + it("hydrates to undefined when nothing carries params, so an untouched save stays byte-identical", () => { + expect( + hydrateTierModelParams({ SIMPLE: ["mini"], MEDIUM: [{ model_name: "opus", litellm_params: {} }] }, undefined), + ).toBeUndefined(); + }); +}); + +describe("serializeTierModelConfigs", () => { + const tiers: ComplexityTiers = { SIMPLE: ["mini"], MEDIUM: ["opus"], COMPLEX: ["opus"], REASONING: [] }; + + it("emits the sibling wire shape per tier and model", () => { + expect( + serializeTierModelConfigs(tiers, { + MEDIUM: { opus: { reasoning_effort: "medium" } }, + COMPLEX: { opus: { reasoning_effort: "high" } }, + }), + ).toEqual({ + MEDIUM: [{ model_name: "opus", litellm_params: { reasoning_effort: "medium" } }], + COMPLEX: [{ model_name: "opus", litellm_params: { reasoning_effort: "high" } }], + }); + }); + + it("prunes params for a model no longer selected in the tier", () => { + expect( + serializeTierModelConfigs(tiers, { MEDIUM: { "removed-model": { reasoning_effort: "low" } } }), + ).toBeUndefined(); + }); + + // Params authored in config.yaml alongside reasoning_effort must survive an edit round-trip. + it("carries params keys this editor has no control for", () => { + expect( + serializeTierModelConfigs(tiers, { MEDIUM: { opus: { reasoning_effort: "medium", max_tokens: 512 } } }), + ).toEqual({ + MEDIUM: [{ model_name: "opus", litellm_params: { reasoning_effort: "medium", max_tokens: 512 } }], + }); + }); + + // This modal renders only the four built-in tiers; params stored under an operator-defined tier + // must pass through rather than being dropped the moment the key became managed. + it("passes tiers this editor does not render through untouched", () => { + expect(serializeTierModelConfigs(tiers, { DEEP_RESEARCH: { opus: { reasoning_effort: "xhigh" } } })).toEqual({ + DEEP_RESEARCH: [{ model_name: "opus", litellm_params: { reasoning_effort: "xhigh" } }], + }); + }); + + it("round-trips what hydration produced", () => { + const stored = { MEDIUM: [{ model_name: "opus", litellm_params: { reasoning_effort: "medium" } }] }; + expect(serializeTierModelConfigs(tiers, hydrateTierModelParams(tiers, stored))).toEqual(stored); + }); + + it("serializes to undefined when nothing is set", () => { + expect(serializeTierModelConfigs(tiers, undefined)).toBeUndefined(); + expect(serializeTierModelConfigs(tiers, { MEDIUM: {} })).toBeUndefined(); + }); +}); + +describe("setTierModelReasoningEffort", () => { + it("sets an effort for a tier and model", () => { + expect(setTierModelReasoningEffort(undefined, "MEDIUM", "opus", "medium")).toEqual({ + MEDIUM: { opus: { reasoning_effort: "medium" } }, + }); + }); + + it("unsetting removes the key and collapses empties back to undefined", () => { + const set = setTierModelReasoningEffort(undefined, "MEDIUM", "opus", "medium"); + expect(setTierModelReasoningEffort(set, "MEDIUM", "opus", undefined)).toBeUndefined(); + }); + + it("unsetting the effort keeps params keys it does not own", () => { + expect( + setTierModelReasoningEffort( + { MEDIUM: { opus: { reasoning_effort: "medium", max_tokens: 512 } } }, + "MEDIUM", + "opus", + undefined, + ), + ).toEqual({ MEDIUM: { opus: { max_tokens: 512 } } }); + }); + + it("leaves other tiers and models alone", () => { + expect( + setTierModelReasoningEffort({ COMPLEX: { opus: { reasoning_effort: "high" } } }, "MEDIUM", "opus", "low"), + ).toEqual({ + COMPLEX: { opus: { reasoning_effort: "high" } }, + MEDIUM: { opus: { reasoning_effort: "low" } }, + }); + }); +}); + +describe("pruneTierModelParams", () => { + it("drops params for models deselected from the tier", () => { + expect( + pruneTierModelParams({ MEDIUM: { opus: { reasoning_effort: "medium" } } }, "MEDIUM", ["mini"]), + ).toBeUndefined(); + }); + + it("keeps params for models still selected", () => { + const current = { MEDIUM: { opus: { reasoning_effort: "medium" } } }; + expect(pruneTierModelParams(current, "MEDIUM", ["opus", "mini"])).toEqual(current); + }); + + it("returns the input unchanged when the tier holds no params", () => { + const current = { COMPLEX: { opus: { reasoning_effort: "high" } } }; + expect(pruneTierModelParams(current, "MEDIUM", [])).toBe(current); + }); +}); diff --git a/ui/litellm-dashboard/src/components/add_model/complexity_router_tiers.ts b/ui/litellm-dashboard/src/components/add_model/complexity_router_tiers.ts index 6f5eb7f877b..cc12746a755 100644 --- a/ui/litellm-dashboard/src/components/add_model/complexity_router_tiers.ts +++ b/ui/litellm-dashboard/src/components/add_model/complexity_router_tiers.ts @@ -1,19 +1,119 @@ import type { ComplexityTiers } from "./ComplexityRouterConfig"; import type { ComplexityTier } from "./KeywordTierRules"; +export type TierModelParams = Record; + +export type TierModelParamsByTier = Record>; + +export const REASONING_EFFORT_OPTIONS = ["none", "minimal", "low", "medium", "high", "xhigh"] as const; +export type ReasoningEffort = (typeof REASONING_EFFORT_OPTIONS)[number]; + +const asRecord = (raw: unknown): Record | undefined => + typeof raw === "object" && raw !== null && !Array.isArray(raw) ? (raw as Record) : undefined; + +const asTierEntryObject = (entry: unknown): { model_name: string; litellm_params: TierModelParams } | undefined => { + const record = asRecord(entry); + if (record === undefined || typeof record.model_name !== "string" || !record.model_name) return undefined; + return { model_name: record.model_name, litellm_params: asRecord(record.litellm_params) ?? {} }; +}; + /** - * A complexity tier maps to `str | list[str]` on the backend - * (litellm/router_strategy/complexity_router/config.py: "string = pin; list = random pick"), - * and the router widens the bare string with `models if isinstance(models, list) else [models]`. + * A complexity tier maps to `str | object | list[str | object]` on the backend + * (litellm/router_strategy/complexity_router/config.py: string/object = pin; list = random pick; + * an object is `{model_name, litellm_params}`), and the router widens a bare value to a list. * * Every UI reader of a STORED complexity_router_config must widen the same way, so this is the * single owner of that rule. Readers of in-memory ComplexityTiers state are already string[] * and do not need it. */ export const normalizeTierModels = (value: unknown): string[] => { - if (Array.isArray(value)) return value.filter((model): model is string => typeof model === "string"); - if (typeof value === "string" && value) return [value]; - return []; + const entries = Array.isArray(value) ? value : [value]; + return entries.flatMap((entry) => { + if (typeof entry === "string" && entry) return [entry]; + const parsed = asTierEntryObject(entry); + return parsed ? [parsed.model_name] : []; + }); +}; + +const tierEntriesWithParams = (value: unknown): [string, TierModelParams][] => + (Array.isArray(value) ? value : [value]) + .map(asTierEntryObject) + .filter((entry): entry is { model_name: string; litellm_params: TierModelParams } => entry !== undefined) + .filter((entry) => Object.keys(entry.litellm_params).length > 0) + .map((entry) => [entry.model_name, entry.litellm_params]); + +/** + * Params can be stored two ways: inline object entries in `tiers`, or the sibling + * `tier_model_configs` key. The backend merges them with `tier_model_configs` winning per + * (tier, model) (config.py `_normalize_tier_model_configs`), so hydration must too. + */ +export const hydrateTierModelParams = ( + storedTiers: unknown, + storedTierModelConfigs: unknown, +): TierModelParamsByTier | undefined => { + const fromInline = Object.entries(asRecord(storedTiers) ?? {}).map( + ([tier, value]) => [tier, tierEntriesWithParams(value)] as const, + ); + const fromSibling = Object.entries(asRecord(storedTierModelConfigs) ?? {}).map( + ([tier, value]) => [tier, tierEntriesWithParams(value)] as const, + ); + const merged = [...fromInline, ...fromSibling].reduce( + (byTier, [tier, entries]) => + entries.length === 0 ? byTier : { ...byTier, [tier]: { ...byTier[tier], ...Object.fromEntries(entries) } }, + {}, + ); + return Object.keys(merged).length > 0 ? merged : undefined; +}; + +/** + * Undefined when empty rather than `{}`, so an untouched router keeps the key out of its payload; + * tiers this editor does not render pass through rather than being dropped now the key is managed. + */ +export const serializeTierModelConfigs = ( + tiers: ComplexityTiers, + tierModelParams: TierModelParamsByTier | undefined, +): Record | undefined => { + if (tierModelParams === undefined) return undefined; + const serialized = Object.entries(tierModelParams) + .map(([tier, byModel]) => { + const selected = (TIER_ORDER as string[]).includes(tier) ? new Set(tiers[tier as ComplexityTier]) : undefined; + const entries = Object.entries(byModel) + .filter(([model, params]) => (selected === undefined || selected.has(model)) && Object.keys(params).length > 0) + .map(([model_name, litellm_params]) => ({ model_name, litellm_params })); + return [tier, entries] as const; + }) + .filter(([, entries]) => entries.length > 0); + return serialized.length > 0 ? Object.fromEntries(serialized) : undefined; +}; + +export const setTierModelReasoningEffort = ( + current: TierModelParamsByTier | undefined, + tier: string, + model: string, + effort: ReasoningEffort | undefined, +): TierModelParamsByTier | undefined => { + const { reasoning_effort: _dropped, ...rest } = current?.[tier]?.[model] ?? {}; + const params = effort === undefined ? rest : { ...rest, reasoning_effort: effort }; + const byModel = Object.fromEntries( + Object.entries({ ...current?.[tier], [model]: params }).filter(([, value]) => Object.keys(value).length > 0), + ); + const next = Object.fromEntries( + Object.entries({ ...current, [tier]: byModel }).filter(([, value]) => Object.keys(value).length > 0), + ); + return Object.keys(next).length > 0 ? next : undefined; +}; + +export const pruneTierModelParams = ( + current: TierModelParamsByTier | undefined, + tier: string, + selectedModels: string[], +): TierModelParamsByTier | undefined => { + if (current?.[tier] === undefined) return current; + const byModel = Object.fromEntries(Object.entries(current[tier]).filter(([model]) => selectedModels.includes(model))); + const next = Object.fromEntries( + Object.entries({ ...current, [tier]: byModel }).filter(([, value]) => Object.keys(value).length > 0), + ); + return Object.keys(next).length > 0 ? next : undefined; }; /** diff --git a/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.test.tsx b/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.test.tsx index 967fc9c458a..07ec36639b4 100644 --- a/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.test.tsx +++ b/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.test.tsx @@ -1,8 +1,28 @@ import { render, screen } from "@testing-library/react"; +import React, { useEffect, useRef } from "react"; +import { useFormContext, useWatch } from "react-hook-form"; import { describe, expect, it } from "vitest"; import { MountedFormHost } from "../../../tests/mounted-form-host"; +import type { MountedFormValues } from "../common_components/MountedFormField"; import ConditionalPublicModelName from "./conditional_public_model_name"; +const WRITE_BUDGET = 20; + +const LoopGuard: React.FC = () => { + const form = useFormContext(); + const mappings = useWatch({ control: form.control, name: "model_mappings" }); + const writes = useRef(0); + + useEffect(() => { + writes.current += 1; + if (writes.current > WRITE_BUDGET) { + throw new Error(`model_mappings changed ${WRITE_BUDGET}+ times: the mapping effects are looping`); + } + }, [mappings]); + + return null; +}; + describe("ConditionalPublicModelName", () => { it("should render", () => { render( @@ -25,4 +45,28 @@ describe("ConditionalPublicModelName", () => { expect(screen.getByText("Public Model Name")).toBeInTheDocument(); expect(screen.getByText("LiteLLM Model Name")).toBeInTheDocument(); }); + + it("settles after rewriting the custom placeholder mapping to the entered model name", () => { + render( + + + + , + ); + + expect(screen.getByDisplayValue("my-custom-model")).toBeInTheDocument(); + expect(screen.getByText("my-custom-model")).toBeInTheDocument(); + expect(screen.queryByDisplayValue("custom")).not.toBeInTheDocument(); + }); }); diff --git a/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.tsx b/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.tsx index 3dc2d781922..b9c128a3d51 100644 --- a/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.tsx +++ b/ui/litellm-dashboard/src/components/add_model/conditional_public_model_name.tsx @@ -1,4 +1,4 @@ -import React, { useEffect, useState } from "react"; +import React, { useEffect, useMemo } from "react"; import type { ColumnDef } from "@tanstack/react-table"; import { useFormContext, useWatch } from "react-hook-form"; import { DataTable } from "@/components/shared/DataTable"; @@ -13,6 +13,13 @@ interface ModelMapping { litellm_model: string; } +const sameMappings = (left: readonly ModelMapping[], right: readonly ModelMapping[]): boolean => + left.length === right.length && + left.every( + (mapping, index) => + mapping.public_name === right[index].public_name && mapping.litellm_model === right[index].litellm_model, + ); + const modelMappingsRule = { validator: async (_: unknown, value: unknown) => { if (!value || (value as ModelMapping[]).length === 0) { @@ -29,15 +36,14 @@ const modelMappingsRule = { const ConditionalPublicModelName: React.FC = () => { const form = useFormContext(); - const [tableKey, setTableKey] = useState(0); // Add a key to force table re-render - // Watch the 'model' field for changes and ensure it's always an array const modelValue = useWatch({ control: form.control, name: "model" }) || []; - const selectedModels = Array.isArray(modelValue) ? modelValue : [modelValue]; + const selectionKey = JSON.stringify(Array.isArray(modelValue) ? modelValue : [modelValue]); + const selectedModels = useMemo(() => JSON.parse(selectionKey) as string[], [selectionKey]); const customModelName = useWatch({ control: form.control, name: "custom_model_name" }) as string | undefined; const showPublicModelName = !selectedModels.includes("all-wildcard"); const selectedProvider = useWatch({ control: form.control, name: "custom_llm_provider" }); - // Force table to re-render when custom model name changes + useEffect(() => { if (customModelName && selectedModels.includes("custom")) { const currentMappings = (form.getValues("model_mappings") as ModelMapping[]) || []; @@ -56,8 +62,9 @@ const ConditionalPublicModelName: React.FC = () => { } return mapping; }); - form.setValue("model_mappings", updatedMappings); - setTableKey((prev) => prev + 1); // Force table re-render + if (!sameMappings(currentMappings, updatedMappings)) { + form.setValue("model_mappings", updatedMappings); + } } }, [customModelName, selectedModels, selectedProvider, form]); @@ -109,7 +116,6 @@ const ConditionalPublicModelName: React.FC = () => { }); form.setValue("model_mappings", mappings); - setTableKey((prev) => prev + 1); // Force table re-render } } }, [selectedModels, customModelName, selectedProvider, form]); @@ -210,7 +216,6 @@ const ConditionalPublicModelName: React.FC = () => { > {(control) => ( row.litellm_model} diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts index 0ae782dc5fc..59254dbbe6f 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts @@ -373,3 +373,59 @@ describe("buildUpdatedComplexityRouterConfig plan-mode minimum tier", () => { expect(result.plan_mode_min_tier).toBe("MEDIUM"); }); }); + +describe("buildUpdatedComplexityRouterConfig tier model params", () => { + const storedWithParams = { + ...STORED, + tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: ["opus"], COMPLEX: ["opus"], REASONING: [] }, + tier_model_configs: { + MEDIUM: [{ model_name: "opus", litellm_params: { reasoning_effort: "medium" } }], + COMPLEX: [{ model_name: "opus", litellm_params: { reasoning_effort: "high" } }], + }, + }; + const formValueWithParams = { + ...FORM_VALUE, + tiers: storedWithParams.tiers, + tier_model_params: { + MEDIUM: { opus: { reasoning_effort: "medium" } }, + COMPLEX: { opus: { reasoning_effort: "high" } }, + }, + }; + + it("round-trips hydrated params on an untouched save", () => { + const result = buildUpdatedComplexityRouterConfig(storedWithParams, formValueWithParams, undefined, hydratedState); + expect(result.tier_model_configs).toEqual(storedWithParams.tier_model_configs); + }); + + // tier_model_configs is managed now that this modal renders a control for it. Before that, the + // stale stored key was carried through, so clearing the last effort could never persist. + it("drops the stored key entirely when the operator unsets every effort", () => { + const result = buildUpdatedComplexityRouterConfig( + storedWithParams, + { ...formValueWithParams, tier_model_params: undefined }, + undefined, + hydratedState, + ); + expect(result).not.toHaveProperty("tier_model_configs"); + }); + + it("drops params for a model removed from its tier", () => { + const result = buildUpdatedComplexityRouterConfig( + storedWithParams, + { + ...formValueWithParams, + tiers: { ...storedWithParams.tiers, COMPLEX: ["gpt-4o-mini"] }, + }, + undefined, + hydratedState, + ); + expect(result.tier_model_configs).toEqual({ + MEDIUM: [{ model_name: "opus", litellm_params: { reasoning_effort: "medium" } }], + }); + }); + + it("emits no tier_model_configs for a config that never had params", () => { + const result = buildUpdatedComplexityRouterConfig(STORED, FORM_VALUE, undefined, hydratedState); + expect(result).not.toHaveProperty("tier_model_configs"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx index 012624454d3..811ff69642a 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx @@ -14,7 +14,12 @@ import ModelChoiceCombobox, { type ModelChoice } from "../add_model/ModelChoiceC import { modelAvailableCall, modelPatchUpdateCall } from "../networking"; import { fetchAvailableModels, ModelGroup } from "@/components/llm_calls/fetch_models"; import RouterConfigBuilder from "../add_model/RouterConfigBuilder"; -import { normalizeTierModels, resolveComplexityDefaultModel } from "../add_model/complexity_router_tiers"; +import { + hydrateTierModelParams, + normalizeTierModels, + resolveComplexityDefaultModel, + serializeTierModelConfigs, +} from "../add_model/complexity_router_tiers"; import { isComplexityRouter } from "../add_model/auto_router_strategies"; import { getKeywordTierRulesError, @@ -66,6 +71,7 @@ interface EditAutoRouterModalProps { // actually renders a control that can set it. const MANAGED_COMPLEXITY_ROUTER_KEYS = new Set([ "tiers", + "tier_model_configs", "default_model", "plan_mode_min_tier", "tier_labels", @@ -149,9 +155,12 @@ export const buildUpdatedComplexityRouterConfig = ( const serializedTierLabels = serializeTierLabels(value.tier_labels); const scorerRuns = heuristicScoringRole(value) !== "never"; + const serializedTierModelConfigs = serializeTierModelConfigs(value.tiers, value.tier_model_params); + return { ...preservedConfig, tiers: value.tiers, + ...(serializedTierModelConfigs && { tier_model_configs: serializedTierModelConfigs }), ...(value.default_model?.trim() && { default_model: value.default_model }), ...(value.plan_mode_min_tier?.trim() && { plan_mode_min_tier: value.plan_mode_min_tier }), ...(serializedTierLabels && { tier_labels: serializedTierLabels }), @@ -342,6 +351,7 @@ const EditAutoRouterModal: React.FC = ({ const hydratedComplexityRouterConfig: ComplexityRouterConfigValue = { tiers: hydratedTiers, + tier_model_params: hydrateTierModelParams(parsedConfig.tiers, parsedConfig.tier_model_configs), default_model: hydratePinnedDefaultModel( parsedConfig.default_model, modelData.litellm_params?.complexity_router_default_model, diff --git a/ui/litellm-dashboard/src/components/llm_calls/fetch_models.tsx b/ui/litellm-dashboard/src/components/llm_calls/fetch_models.tsx index 9fbaa868998..4f4c901a1c2 100644 --- a/ui/litellm-dashboard/src/components/llm_calls/fetch_models.tsx +++ b/ui/litellm-dashboard/src/components/llm_calls/fetch_models.tsx @@ -6,6 +6,7 @@ import { modelAvailableCall, modelHubCall } from "@/components/networking"; export interface ModelGroup { model_group: string; mode?: string; + supports_reasoning?: boolean; } interface AvailableModel { @@ -13,6 +14,7 @@ interface AvailableModel { model_name?: string | null; id?: string | null; mode?: string | null; + supports_reasoning?: boolean | null; } export const fetchAvailableModelsForTeam = async (accessToken: string, teamId: string): Promise => { @@ -36,6 +38,7 @@ export const fetchAvailableModels = async (accessToken: string): Promise ({ model_group: item.model_group || item.id || item.model_name || "", mode: item.mode || undefined, + supports_reasoning: item.supports_reasoning === true || undefined, })) .filter((model: ModelGroup) => model.model_group !== ""); diff --git a/ui/litellm-dashboard/src/components/shared/PaginationStatusAlerts.test.tsx b/ui/litellm-dashboard/src/components/shared/PaginationStatusAlerts.test.tsx new file mode 100644 index 00000000000..3698b68155e --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/PaginationStatusAlerts.test.tsx @@ -0,0 +1,62 @@ +import { fireEvent, render } from "@testing-library/react"; +import { describe, expect, it, vi } from "vitest"; + +import PaginationStatusAlerts from "./PaginationStatusAlerts"; + +describe("PaginationStatusAlerts", () => { + it("shows page progress and wires the Stop button while fetching", () => { + const cancel = vi.fn(); + const { getByRole, getByText } = render( + , + ); + + expect(getByText(/Currently fetching spend data: fetched 7 \/ 42 pages/)).toBeInTheDocument(); + fireEvent.click(getByRole("button", { name: "Stop" })); + expect(cancel).toHaveBeenCalledTimes(1); + }); + + it("shows the partial-data notice after a cancel, frozen at the last fetched page", () => { + const { getByText } = render( + , + ); + + expect(getByText("Showing partial spend data (7/42 pages loaded)")).toBeInTheDocument(); + }); + + it("names the subject it is fetching", () => { + const { getByText } = render( + , + ); + + expect(getByText(/Currently fetching agent data: fetched 1 \/ 3 pages/)).toBeInTheDocument(); + }); + + it("renders nothing when idle", () => { + const { container } = render( + , + ); + + expect(container).toBeEmptyDOMElement(); + }); +}); diff --git a/ui/litellm-dashboard/src/components/shared/PaginationStatusAlerts.tsx b/ui/litellm-dashboard/src/components/shared/PaginationStatusAlerts.tsx new file mode 100644 index 00000000000..af8b8439703 --- /dev/null +++ b/ui/litellm-dashboard/src/components/shared/PaginationStatusAlerts.tsx @@ -0,0 +1,51 @@ +import { ExternalLink, Loader2 } from "lucide-react"; + +import { Alert, AlertDescription } from "@/components/shared/Alert"; +import { Button } from "@/components/ui/button"; + +interface PaginationStatusAlertsProps { + isFetchingMore: boolean; + cancelled: boolean; + progress: { currentPage: number; totalPages: number }; + cancel: () => void; + subject?: string; +} + +const PaginationStatusAlerts = ({ + isFetchingMore, + cancelled, + progress, + cancel, + subject = "spend data", +}: PaginationStatusAlertsProps) => ( + <> + {isFetchingMore && ( + + + + + Currently fetching {subject}: fetched {progress.currentPage} / {progress.totalPages} pages. Charts will + update periodically as data loads. Moving off of this page will stop and reset this. To continue using the + UI in the meantime,{" "} + + open a new tab + + . + + + + + )} + {cancelled && ( + + + Showing partial {subject} ({progress.currentPage}/{progress.totalPages} pages loaded) + + + )} + +); + +export default PaginationStatusAlerts; diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 790ebcdd227..23b34ba2726 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -849,11 +849,12 @@ export interface paths { * serve and duplicates them against baseline_model. A key can hold one active job per * direction, so both questions can run at once. * - * Shadow responses are never served to users. Each key samples until it has judged - * max_turns turns of its own traffic, the job's window ends, or the job is stopped, so one - * key running out of budget does not end sampling for the others; sampling changes - * propagate to pods within about 10 seconds. Shadow and judge calls bill to the shadowed - * key but are excluded from request counts and auto-router adoption metrics. + * Shadow responses are never served to users. Each key samples until its recorded eval + * spend, the shadow and judge calls' own cost, reaches max_budget dollars, the job's + * window ends, or the job is stopped, so one key running out of budget does not end + * sampling for the others; sampling changes propagate to pods within about 10 seconds. + * Shadow and judge calls bill to the shadowed key but are excluded from request counts + * and auto-router adoption metrics. */ post: operations["start_shadow_eval_auto_router_shadow_eval_start_post"]; delete?: never; @@ -24134,6 +24135,11 @@ export interface components { * @description Maximum retention period for auto-router benchmark session rollup rows (e.g., '365d'). Rows whose last turn is older than this are deleted by the spend log cleanup job, on that job's schedule. Unset means rollup rows are never deleted. */ maximum_autorouter_session_retention_period?: string | null; + /** + * Maximum Health Check Retention Period + * @description Maximum retention period for health-check rows (e.g., '30d'). Rows whose checked_at is older than this are deleted by the spend log cleanup job, on that job's schedule. Unset means rows are never deleted. Set this well above health_check_interval because /health and the UI read the latest row per model. + */ + maximum_health_check_retention_period?: string | null; /** * Maximum Spend Logs Cleanup Batch Size * @description Rows deleted per DELETE statement by the spend log cleanup job. Defaults to 1000. @@ -33305,11 +33311,21 @@ export interface components { * @description Masked display name (sk-...) of the shadowed key, resolved at read time like key_alias */ key_name?: string | null; + /** + * Max Budget + * @description This key's own USD budget for the eval's shadow and judge spend, independent of its siblings'; None on jobs created before spend budgets existed, which max_turns alone bounds + */ + max_budget?: number | null; /** * Max Turns - * @description This key's own sample budget, independent of its siblings' + * @description This key's sample-count ceiling: the whole budget for jobs created before max_budget existed, and the error-loop safety valve otherwise */ max_turns: number; + /** + * Spend + * @description This key's recorded shadow plus judge spend in USD, the same figure the sampler budgets against max_budget; populated on list and detail responses and frozen at stopped_at exactly like attempt_count + */ + spend?: number | null; /** * Stopped At * @description When this key's slot was stamped free, whether its own budget ran out, the window closed, or an operator stopped the job; status is derived, so a spent budget reads completed even while this is still unset @@ -33626,7 +33642,7 @@ export interface components { StartShadowEvalRequest: { /** * Api Key Ids - * @description The hashed virtual keys whose traffic will be shadowed. Shadow evaluation runs ONLY on these keys' traffic; requests made with any other key are not sampled. Each key carries its own max_turns budget, so one key exhausting its budget leaves the others sampling. At most 100 keys per job, which also bounds every read the job's endpoints make. + * @description The hashed virtual keys whose traffic will be shadowed. Shadow evaluation runs ONLY on these keys' traffic; requests made with any other key are not sampled. Each key carries its own max_budget spend budget, so one key exhausting its budget leaves the others sampling. At most 100 keys per job, which also bounds every read the job's endpoints make. */ api_key_ids: string[]; /** @@ -33654,11 +33670,11 @@ export interface components { */ judge_model: string; /** - * Max Turns - * @description Per-key sample budget: the job judges at most this many turns of EACH scoped key's traffic, so a job over N keys judges at most N times max_turns turns. This is also the spend bound; expected judge cost is roughly that turn ceiling times one judge call - * @default 200 + * Max Budget + * @description Per-key USD budget for the eval's own overhead, the shadow-arm and judge calls, priced with the same figures the spend pipeline bills. EACH scoped key samples until its recorded eval spend reaches this, so a job over N keys spends at most about N times max_budget; in-flight samples can overshoot the cap by one sampling cache window + * @default 10 */ - max_turns: number; + max_budget: number; /** * Router Name * @description The auto-router under evaluation, in either direction diff --git a/uv.lock b/uv.lock index 53bb0cb8f82..6b18be68c92 100644 --- a/uv.lock +++ b/uv.lock @@ -10,7 +10,7 @@ resolution-markers = [ ] [options] -exclude-newer = "2026-08-17T01:06:38.502388Z" +exclude-newer = "2026-08-17T21:26:36.028845Z" exclude-newer-span = "P3D" [manifest] @@ -4661,12 +4661,12 @@ proxy-dev = [ [[package]] name = "litellm-enterprise" -version = "0.1.57" +version = "0.1.58" source = { editable = "enterprise" } [[package]] name = "litellm-proxy-extras" -version = "0.4.87" +version = "0.4.88" source = { editable = "litellm-proxy-extras" } [[package]]