Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_lit5879_semantic_cache_embedding_timeout

This commit is contained in:
mateo-berri 2026-08-20 16:59:57 -07:00
commit 7d23d41cc4
282 changed files with 20365 additions and 3882 deletions

View file

@ -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

View file

@ -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(

View file

@ -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())

View file

@ -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:?} \

View file

@ -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

View file

@ -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"

View file

@ -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

View file

@ -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

View file

@ -105,13 +105,13 @@
"limit": 109
},
"reportUnknownMemberType": {
"limit": 39017
"limit": 39011
},
"reportUnknownParameterType": {
"limit": 19885
},
"reportUnknownVariableType": {
"limit": 30572
"limit": 30569
},
"reportUnnecessaryCast": {
"limit": 117

View file

@ -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:

View file

@ -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==",

View file

@ -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:

View file

@ -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:

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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;

View file

@ -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())

View file

@ -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==",

View file

@ -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",
}
}

View file

@ -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",
}
}

View file

@ -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,

View file

@ -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<reqwest::Client> = 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())
})
}

View file

@ -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<Map<String, Value>>,
) -> CoreResult<Vec<(String, String)>> {
shared_string_headers(HEADER_CONTEXT, extra_headers)
}

View file

@ -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<String>,
}
#[derive(Clone, Debug, Default, PartialEq, Eq)]
pub struct Conversation {
pub system: Vec<String>,
pub turns: Vec<Turn>,
}
/// 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<String> {
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::<Turn>::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<ChatMessage> {
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()]);
}
}

View file

@ -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<ChatCompletionsResponse> {
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<Vec<(String, String)>> {
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<String, String> = 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<Vec<(String, String)>> {
match &request.auth {
ChatCompletionsAuth::AwsSigV4 { .. } => Err(CoreError::Unsupported(
"AWS SigV4 requires the bedrock-auth feature",
)),
_ => Ok(request.upstream_headers.clone()),
}
}

View file

@ -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<ChatCompletionsResponse> {
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<String, Value>,
) -> 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;

View file

@ -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<Vec<ChatMessage>> {
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<ProviderChatCompletionsRequest> {
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,
})
}

View file

@ -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);
}
}

View file

@ -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, &params)
}
#[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::<usize>().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<String>) {
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, .. }
));
}
}

View file

@ -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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String>;
fn auth(
&self,
api_key: Option<&str>,
model: &str,
optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<ChatCompletionsAuth>;
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<String, Value>,
) -> Option<Unsupported> {
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<ChatMessage>,
optional_params: Map<String, Value>,
) -> CoreResult<ProviderChatRequestData>;
fn transform_response(
&self,
model: &str,
response: ProviderChatResponseData,
) -> CoreResult<ChatCompletionsResponse>;
}
pub fn unsupported_param(
supported: &'static [&'static str],
config: &'static [&'static str],
optional_params: &Map<String, Value>,
) -> Option<Unsupported> {
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<Unsupported> {
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")),
}
}

View file

@ -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<String, Value>,
pub api_key: Option<&'a str>,
pub api_base: Option<&'a str>,
pub custom_llm_provider: Option<&'a str>,
pub extra_headers: Option<Map<String, Value>>,
pub timeout: Option<Duration>,
}
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<String, Value>,
pub(super) timeout: Option<Duration>,
}
/// 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<Value>),
}
#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)]
pub struct ChatMessage {
pub role: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub content: Option<ChatMessageContent>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub name: Option<String>,
#[serde(flatten)]
pub extra: Map<String, Value>,
}
/// 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<String>,
}
#[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<ChatCompletionsChoice>,
pub usage: ChatCompletionsUsage,
}

View file

@ -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]";

View file

@ -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 {

View file

@ -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<Map<String, Value>>,
) -> CoreResult<Vec<(String, String)>> {
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()
)]));
}
}

View file

@ -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;

View file

@ -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<Map<String, Value>>,
) -> CoreResult<Vec<(String, String)>> {
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)
}

View file

@ -0,0 +1 @@
pub mod transformation;

View file

@ -0,0 +1,444 @@
use super::*;
use serde_json::json;
fn messages(value: Value) -> Vec<ChatMessage> {
serde_json::from_value(value).expect("valid messages")
}
fn params(value: Value) -> Map<String, Value> {
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<ChatCompletionsResponse> {
ANTHROPIC_CHAT_COMPLETIONS_CONFIG
.transform_response("claude-sonnet-4-5", ProviderChatResponseData { body })
}
fn reason(msgs: Value, opts: Value) -> Option<Unsupported> {
ANTHROPIC_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), &params(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"),
]
);
}

View file

@ -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<String, Value>) -> Value {
let messages: Vec<Value> = conversation
.turns
.iter()
.map(|turn| {
json!({
"role": turn.role.as_str(),
"content": turn.texts.iter().map(|text| text_block(text)).collect::<Vec<_>>(),
})
})
.collect();
let system: Vec<Value> = 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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
Ok(complete_anthropic_url(api_base, env_lookup))
}
fn auth(
&self,
api_key: Option<&str>,
_model: &str,
_optional_params: &Map<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<ChatCompletionsAuth> {
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<String, Value>,
) -> Option<Unsupported> {
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<ChatMessage>,
optional_params: Map<String, Value>,
) -> CoreResult<ProviderChatRequestData> {
Ok(ProviderChatRequestData {
body: anthropic_body(model, &build_conversation(&messages), optional_params),
})
}
fn transform_response(
&self,
_model: &str,
response: ProviderChatResponseData,
) -> CoreResult<ChatCompletionsResponse> {
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;

View file

@ -1 +1,2 @@
pub mod chat_completions;
pub mod messages;

View file

@ -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<String>) {
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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> 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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> 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::*;

View file

@ -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<String, String>) -> BTreeMap<String, String> {
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<String>) {
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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> 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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> 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<String, Value>) -> Option<Credentials> {
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();

View file

@ -0,0 +1 @@
pub mod transformation;

View file

@ -0,0 +1,580 @@
use super::*;
use serde_json::json;
fn messages(value: Value) -> Vec<ChatMessage> {
serde_json::from_value(value).expect("valid messages")
}
fn params(value: Value) -> Map<String, Value> {
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<ChatCompletionsResponse> {
BEDROCK_CHAT_COMPLETIONS_CONFIG.transform_response(
"anthropic.claude-sonnet-4-5-v1:0",
ProviderChatResponseData { body },
)
}
fn reason(msgs: Value, opts: Value) -> Option<Unsupported> {
BEDROCK_CHAT_COMPLETIONS_CONFIG.unsupported_reason(&messages(msgs), &params(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<String>| {
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(&params(json!({"aws_access_key_id": "AKIA"}))).is_none());
assert!(
host_supplied_credentials(&params(
json!({"aws_access_key_id": " ", "aws_secret_access_key": "s"})
))
.is_none()
);
assert!(host_supplied_credentials(&Map::new()).is_none());
}

View file

@ -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<String, Value>) -> Value {
let messages: Vec<Value> = conversation
.turns
.iter()
.map(|turn| {
json!({
"role": turn.role.as_str(),
"content": turn.texts.iter().map(|text| json!({"text": text})).collect::<Vec<_>>(),
})
})
.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<Value> = 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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<String> {
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}", &region));
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<String, Value>,
env_lookup: &dyn Fn(&str) -> Option<String>,
) -> CoreResult<ChatCompletionsAuth> {
// 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<String, Value>,
) -> Option<Unsupported> {
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<ChatMessage>,
optional_params: Map<String, Value>,
) -> CoreResult<ProviderChatRequestData> {
Ok(ProviderChatRequestData {
body: converse_body(&build_conversation(&messages), &optional_params),
})
}
fn transform_response(
&self,
model: &str,
response: ProviderChatResponseData,
) -> CoreResult<ChatCompletionsResponse> {
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;

View file

@ -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";

View file

@ -5,4 +5,5 @@
#[cfg(feature = "bedrock-auth")]
pub mod audio_transcription;
pub mod aws_base;
pub mod chat_completions;
mod constants;

View file

@ -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<Map<String, Value>>,
@ -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<Py<PyAny>> {
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<String, Value>,
Option<Map<String, Value>>,
Option<Duration>,
);
fn marshal_chat_completions_inputs(
py: Python<'_>,
messages: Py<PyAny>,
optional_params: Option<Py<PyAny>>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<MarshaledChatCompletionsInputs> {
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<PyAny>,
optional_params: Option<Py<PyAny>>,
custom_llm_provider: Option<String>,
) -> PyResult<Option<String>> {
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<PyAny>,
optional_params: Option<Py<PyAny>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Py<PyAny>> {
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<PyAny>,
optional_params: Option<Py<PyAny>>,
api_key: Option<String>,
api_base: Option<String>,
custom_llm_provider: Option<String>,
extra_headers: Option<Py<PyAny>>,
timeout_seconds: Option<f64>,
) -> PyResult<Bound<'_, PyAny>> {
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<Py<PyAny>> {
let stats = PyDict::new(py);
@ -439,12 +630,18 @@ fn gil_stats(py: Python<'_>) -> PyResult<Py<PyAny>> {
#[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::<RustBridgeDeclined>())?;
module.add("RustUpstreamError", py.get_type::<RustUpstreamError>())?;
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::<ResponsesWebSocketConnection>()?;
module.add_function(wrap_pyfunction!(gil_stats, module)?)?;
Ok(())

View file

@ -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"""

View file

@ -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,

View file

@ -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 ###########################
########################################################################################

View file

@ -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 ``<scheme> `` removed, or unchanged when absent.
Callers supply both a bare credential and a complete header value, so prefixing
unconditionally yields ``Bearer Bearer <jwt>``. 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 <credentials>`` 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

View file

@ -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],

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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),
)

View file

@ -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.

View file

@ -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

View file

@ -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.

View file

@ -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,
)

View file

@ -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:

View file

@ -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 (

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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,

View file

@ -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",

View file

@ -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

View file

@ -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(

View file

@ -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

View file

@ -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: <json>\\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: <json>\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(

View file

@ -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,

View file

@ -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,

View file

@ -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",

View file

@ -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,

View file

@ -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

View file

@ -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)

View file

@ -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

View file

@ -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:

View file

@ -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

View file

@ -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

View file

@ -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", "")

View file

@ -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
),
)

View file

@ -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

View file

@ -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,
)

View file

@ -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))

View file

@ -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,
)

View file

@ -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")
)

View file

@ -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())

View file

@ -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 ###

View file

@ -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",

View file

@ -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())

View file

@ -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:

View file

@ -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(

Some files were not shown because too many files have changed in this diff Show more