mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-08 03:08:45 +00:00
Merge branch 'litellm_internal_staging' of https://github.com/BerriAI/litellm into litellm_lit5879_semantic_cache_embedding_timeout
This commit is contained in:
commit
7d23d41cc4
282 changed files with 20365 additions and 3882 deletions
17
.github/ci-coverage-allowlist.yml
vendored
17
.github/ci-coverage-allowlist.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
39
.github/scripts/assert_ci_coverage.py
vendored
39
.github/scripts/assert_ci_coverage.py
vendored
|
|
@ -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(
|
||||
|
|
|
|||
149
.github/scripts/assert_workflow_dir_hygiene.py
vendored
Normal file
149
.github/scripts/assert_workflow_dir_hygiene.py
vendored
Normal 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())
|
||||
7
.github/workflows/_test-unit-base.yml
vendored
7
.github/workflows/_test-unit-base.yml
vendored
|
|
@ -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:?} \
|
||||
|
|
|
|||
3
.github/workflows/ci-coverage.yml
vendored
3
.github/workflows/ci-coverage.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
7
.github/workflows/test-linting.yml
vendored
7
.github/workflows/test-linting.yml
vendored
|
|
@ -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"
|
||||
|
|
|
|||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
5
Makefile
5
Makefile
|
|
@ -160,6 +160,7 @@ lint-format-check-changed: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
|||
# Linting targets
|
||||
lint-ruff: $(LINT_DEP_INSTALL)
|
||||
cd litellm && $(UV_RUN) ruff check . && cd ..
|
||||
$(UV_RUN) ruff check --config ruff-tests.toml tests
|
||||
|
||||
# faster linter for developing ...
|
||||
# inspiration from:
|
||||
|
|
@ -205,8 +206,8 @@ lint-type-discipline: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
|||
$(UV_RUN) python scripts/type_discipline_gate.py --base origin/litellm_internal_staging
|
||||
|
||||
# Test-quality budget (zero-assert / mock-echo tests, sys.path.insert, raw env writes,
|
||||
# litellm module-global mutation, credential-gated skips), counted across tests/ the
|
||||
# same delta-vs-base way.
|
||||
# litellm module-global mutation, credential-gated skips, conftest snapshot
|
||||
# inventory), counted across tests/ the same delta-vs-base way.
|
||||
lint-test-quality: $(LINT_DEP_INSTALL) $(LINT_DEP_BASE)
|
||||
$(UV_RUN) python scripts/test_quality_gate.py --base origin/litellm_internal_staging
|
||||
|
||||
|
|
|
|||
|
|
@ -105,13 +105,13 @@
|
|||
"limit": 109
|
||||
},
|
||||
"reportUnknownMemberType": {
|
||||
"limit": 39017
|
||||
"limit": 39011
|
||||
},
|
||||
"reportUnknownParameterType": {
|
||||
"limit": 19885
|
||||
},
|
||||
"reportUnknownVariableType": {
|
||||
"limit": 30572
|
||||
"limit": 30569
|
||||
},
|
||||
"reportUnnecessaryCast": {
|
||||
"limit": 117
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
116
helm/litellm/tests/database_auth_tests.yaml
Normal file
116
helm/litellm/tests/database_auth_tests.yaml
Normal 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
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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;
|
||||
|
|
@ -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())
|
||||
|
||||
|
|
|
|||
|
|
@ -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==",
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
15
litellm-rust/crates/core/src/chat_completions/client.rs
Normal file
15
litellm-rust/crates/core/src/chat_completions/client.rs
Normal 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())
|
||||
})
|
||||
}
|
||||
|
|
@ -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)
|
||||
}
|
||||
254
litellm-rust/crates/core/src/chat_completions/conversation.rs
Normal file
254
litellm-rust/crates/core/src/chat_completions/conversation.rs
Normal 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()]);
|
||||
}
|
||||
}
|
||||
147
litellm-rust/crates/core/src/chat_completions/handler.rs
Normal file
147
litellm-rust/crates/core/src/chat_completions/handler.rs
Normal 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()),
|
||||
}
|
||||
}
|
||||
59
litellm-rust/crates/core/src/chat_completions/mod.rs
Normal file
59
litellm-rust/crates/core/src/chat_completions/mod.rs
Normal 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;
|
||||
118
litellm-rust/crates/core/src/chat_completions/prepare.rs
Normal file
118
litellm-rust/crates/core/src/chat_completions/prepare.rs
Normal 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,
|
||||
})
|
||||
}
|
||||
101
litellm-rust/crates/core/src/chat_completions/response_utils.rs
Normal file
101
litellm-rust/crates/core/src/chat_completions/response_utils.rs
Normal 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);
|
||||
}
|
||||
}
|
||||
820
litellm-rust/crates/core/src/chat_completions/tests.rs
Normal file
820
litellm-rust/crates/core/src/chat_completions/tests.rs
Normal 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, ¶ms)
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_accepts_what_prepare_accepts() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_declines_without_resolving_credentials_or_calling_out() {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true}),
|
||||
),
|
||||
Some("streaming")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"openai/gpt-4o",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"claude-sonnet-4-5",
|
||||
None,
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({}),
|
||||
),
|
||||
Some("provider is not on the rust chat completions path")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
json!("nope"),
|
||||
json!({})
|
||||
),
|
||||
Some("unreadable message list")
|
||||
);
|
||||
assert_eq!(
|
||||
decline_reason("anthropic/claude-sonnet-4-5", None, json!([]), json!({})),
|
||||
Some("empty message list")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn the_gate_agrees_with_prepare_on_every_case_it_accepts() {
|
||||
// A gate that accepts what prepare then declines would make the host emit
|
||||
// its pre-call logging on a path that falls back, so pin the agreement.
|
||||
for (messages, params) in [
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 8}),
|
||||
),
|
||||
(
|
||||
json!([{"role": "system", "content": "s"}, {"role": "user", "content": "hi"}]),
|
||||
json!({"temperature": 0.1}),
|
||||
),
|
||||
(
|
||||
json!([{"role": "user", "content": "hi"}, {"role": "assistant", "content": "yo"}]),
|
||||
json!({}),
|
||||
),
|
||||
] {
|
||||
assert_eq!(
|
||||
decline_reason(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
messages.clone(),
|
||||
params.clone()
|
||||
),
|
||||
None,
|
||||
"gate declined {messages}"
|
||||
);
|
||||
prepare_chat_completions_call(request(
|
||||
"anthropic/claude-sonnet-4-5",
|
||||
None,
|
||||
messages.clone(),
|
||||
params,
|
||||
))
|
||||
.unwrap_or_else(|error| panic!("prepare declined {messages}: {error}"));
|
||||
}
|
||||
}
|
||||
|
||||
mod round_trip {
|
||||
use super::*;
|
||||
use tokio::io::{AsyncReadExt, AsyncWriteExt};
|
||||
use tokio::net::{TcpListener, TcpStream};
|
||||
|
||||
use crate::chat_completions::chat_completions;
|
||||
|
||||
async fn read_http_request(socket: &mut TcpStream) -> String {
|
||||
let mut request = Vec::new();
|
||||
let mut buffer = [0_u8; 1024];
|
||||
let header_end = loop {
|
||||
let n = socket.read(&mut buffer).await.expect("reads request");
|
||||
if n == 0 {
|
||||
break request.len();
|
||||
}
|
||||
request.extend_from_slice(&buffer[..n]);
|
||||
if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") {
|
||||
break position + 4;
|
||||
}
|
||||
};
|
||||
let headers = String::from_utf8_lossy(&request[..header_end]);
|
||||
let content_length = headers
|
||||
.lines()
|
||||
.find_map(|line| {
|
||||
let (name, value) = line.split_once(':')?;
|
||||
name.eq_ignore_ascii_case("content-length")
|
||||
.then(|| value.trim().parse::<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, .. }
|
||||
));
|
||||
}
|
||||
}
|
||||
155
litellm-rust/crates/core/src/chat_completions/transformation.rs
Normal file
155
litellm-rust/crates/core/src/chat_completions/transformation.rs
Normal 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")),
|
||||
}
|
||||
}
|
||||
112
litellm-rust/crates/core/src/chat_completions/types.rs
Normal file
112
litellm-rust/crates/core/src/chat_completions/types.rs
Normal 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,
|
||||
}
|
||||
|
|
@ -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]";
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
112
litellm-rust/crates/core/src/http_utils.rs
Normal file
112
litellm-rust/crates/core/src/http_utils.rs
Normal 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()
|
||||
)]));
|
||||
}
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
}
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
pub mod transformation;
|
||||
|
|
@ -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), ¶ms(opts))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_the_messages_body_python_builds() {
|
||||
let body = transform(
|
||||
"claude-sonnet-4-5",
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]),
|
||||
json!({"max_tokens": 128, "temperature": 0.2}),
|
||||
);
|
||||
assert_eq!(
|
||||
body,
|
||||
json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"messages": [
|
||||
{"role": "user", "content": [{"type": "text", "text": "hi"}]}
|
||||
],
|
||||
"system": [{"type": "text", "text": "be terse"}],
|
||||
"max_tokens": 128,
|
||||
"temperature": 0.2
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn omits_system_when_no_system_message_is_present() {
|
||||
let body = transform(
|
||||
"claude-sonnet-4-5",
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"max_tokens": 16}),
|
||||
);
|
||||
assert!(body.get("system").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merges_consecutive_turns_and_wraps_every_text_in_a_block() {
|
||||
let body = transform(
|
||||
"claude-sonnet-4-5",
|
||||
json!([
|
||||
{"role": "user", "content": "one"},
|
||||
{"role": "user", "content": [{"type": "text", "text": "two"}]},
|
||||
{"role": "assistant", "content": "ack"}
|
||||
]),
|
||||
json!({"max_tokens": 16}),
|
||||
);
|
||||
assert_eq!(
|
||||
body["messages"],
|
||||
json!([
|
||||
{"role": "user", "content": [
|
||||
{"type": "text", "text": "one"},
|
||||
{"type": "text", "text": "two"}
|
||||
]},
|
||||
{"role": "assistant", "content": [{"type": "text", "text": "ack"}]}
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn right_strips_a_trailing_assistant_prefill_like_python() {
|
||||
let body = transform(
|
||||
"claude-sonnet-4-5",
|
||||
json!([
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "Argentina "}
|
||||
]),
|
||||
json!({"max_tokens": 16}),
|
||||
);
|
||||
assert_eq!(
|
||||
body["messages"][1]["content"][0]["text"],
|
||||
json!("Argentina")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn passes_every_supported_param_through_untouched() {
|
||||
let body = transform(
|
||||
"claude-sonnet-4-5",
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({
|
||||
"max_tokens": 64,
|
||||
"temperature": 0.1,
|
||||
"top_p": 0.9,
|
||||
"stop_sequences": ["STOP"]
|
||||
}),
|
||||
);
|
||||
assert_eq!(body["max_tokens"], json!(64));
|
||||
assert_eq!(body["temperature"], json!(0.1));
|
||||
assert_eq!(body["top_p"], json!(0.9));
|
||||
assert_eq!(body["stop_sequences"], json!(["STOP"]));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_top_k_because_python_gates_it_by_model_below_this_point() {
|
||||
// `temperature` and `top_p` arrive already resolved, because
|
||||
// `map_openai_params` applies `_apply_sampling_param` to them before the
|
||||
// gate runs. `top_k` bypasses that and is gated inside `transform_request`,
|
||||
// the function this route replaces, so forwarding it would send `top_k` to
|
||||
// a model that removed sampling params and take a 400 after the call, where
|
||||
// Python drops it and succeeds.
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"top_k": 40})
|
||||
),
|
||||
Some(Unsupported("unrecognized request parameter"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_streaming_before_anything_else() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true, "max_tokens": 16})
|
||||
),
|
||||
Some(Unsupported("streaming"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_an_explicit_stream_false() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": false, "max_tokens": 16})
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_any_param_outside_the_allowlist() {
|
||||
for param in [
|
||||
json!({"tools": []}),
|
||||
json!({"tool_choice": {"type": "auto"}}),
|
||||
json!({"thinking": {"type": "enabled"}}),
|
||||
json!({"system": "injected"}),
|
||||
json!({"metadata": {"user_id": "u1"}}),
|
||||
json!({"output_config": {"effort": "high"}}),
|
||||
] {
|
||||
assert_eq!(
|
||||
reason(json!([{"role": "user", "content": "hi"}]), param.clone()),
|
||||
Some(Unsupported("unrecognized request parameter")),
|
||||
"expected {param} to decline"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_tool_calls_tool_results_and_multimodal_content() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": null, "tool_calls": [
|
||||
{"id": "c1", "type": "function",
|
||||
"function": {"name": "f", "arguments": "{}"}}
|
||||
]}
|
||||
]),
|
||||
json!({})
|
||||
),
|
||||
Some(Unsupported("unrecognized message field"))
|
||||
);
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "tool", "tool_call_id": "c1", "content": "ok"}
|
||||
]),
|
||||
json!({})
|
||||
),
|
||||
Some(Unsupported("unrecognized message field"))
|
||||
);
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": [
|
||||
{"type": "image_url", "image_url": {"url": "https://x/y.png"}}
|
||||
]}]),
|
||||
json!({})
|
||||
),
|
||||
Some(Unsupported("non-text message content"))
|
||||
);
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": [
|
||||
{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}}
|
||||
]}]),
|
||||
json!({})
|
||||
),
|
||||
Some(Unsupported("non-text message content"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_a_message_whose_content_list_is_empty() {
|
||||
// An empty list passes every per-part check, so without this it would reach
|
||||
// the provider as an empty `content` array and fail after the call rather
|
||||
// than declining to Python before it.
|
||||
assert_eq!(
|
||||
reason(json!([{"role": "user", "content": []}]), json!({})),
|
||||
Some(Unsupported("message without content"))
|
||||
);
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]),
|
||||
json!({})
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_a_conversation_that_does_not_open_on_a_user_turn() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "assistant", "content": "prefill"}
|
||||
]),
|
||||
json!({})
|
||||
),
|
||||
Some(Unsupported("conversation does not open on a user turn"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_a_plain_text_conversation() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "hello"},
|
||||
{"role": "user", "content": [{"type": "text", "text": "again"}]}
|
||||
]),
|
||||
json!({"max_tokens": 16, "temperature": 0.5})
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn normalizes_a_text_response_into_openai_shape() {
|
||||
let response = transform_response(json!({
|
||||
"id": "msg_123",
|
||||
"type": "message",
|
||||
"role": "assistant",
|
||||
"model": "claude-sonnet-4-5-20260101",
|
||||
"content": [{"type": "text", "text": "hello"}, {"type": "text", "text": " there"}],
|
||||
"stop_reason": "end_turn",
|
||||
"stop_sequence": null,
|
||||
"usage": {"input_tokens": 11, "output_tokens": 4}
|
||||
}))
|
||||
.expect("response transforms");
|
||||
|
||||
assert_eq!(response.model, "claude-sonnet-4-5-20260101");
|
||||
assert_eq!(response.choices.len(), 1);
|
||||
assert_eq!(response.choices[0].index, 0);
|
||||
assert_eq!(response.choices[0].message.role, "assistant");
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("hello there")
|
||||
);
|
||||
assert_eq!(response.choices[0].finish_reason, "stop");
|
||||
assert_eq!(response.usage.prompt_tokens, 11);
|
||||
assert_eq!(response.usage.completion_tokens, 4);
|
||||
assert_eq!(response.usage.total_tokens, 15);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn folds_cache_tokens_into_prompt_tokens_like_python() {
|
||||
let response = transform_response(json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {
|
||||
"input_tokens": 10,
|
||||
"output_tokens": 2,
|
||||
"cache_read_input_tokens": 5,
|
||||
"cache_creation_input_tokens": 3
|
||||
}
|
||||
}))
|
||||
.expect("response transforms");
|
||||
assert_eq!(response.usage.prompt_tokens, 18);
|
||||
assert_eq!(response.usage.total_tokens, 20);
|
||||
assert_eq!(response.usage.prompt_tokens_details.cached_tokens, 5);
|
||||
assert_eq!(
|
||||
response.usage.prompt_tokens_details.cache_creation_tokens,
|
||||
3
|
||||
);
|
||||
assert_eq!(response.usage.prompt_tokens_details.text_tokens, 10);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn maps_max_tokens_stop_reason_to_length() {
|
||||
let response = transform_response(json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
"stop_reason": "max_tokens",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1}
|
||||
}))
|
||||
.expect("response transforms");
|
||||
assert_eq!(response.choices[0].finish_reason, "length");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_refusal_returns_the_completion_python_returns() {
|
||||
// `refusal` is a stop_reason, not a content block type, so the content is
|
||||
// ordinary text and this normalizes rather than declining. Python maps it
|
||||
// to content_filter in _FINISH_REASON_MAP and returns the completion.
|
||||
let response = transform_response(json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [{"type": "text", "text": "I can't help with that."}],
|
||||
"stop_reason": "refusal",
|
||||
"usage": {"input_tokens": 9, "output_tokens": 6}
|
||||
}))
|
||||
.expect("a refusal still transforms");
|
||||
assert_eq!(response.choices[0].finish_reason, "content_filter");
|
||||
assert_eq!(
|
||||
response.choices[0].message.content.as_deref(),
|
||||
Some("I can't help with that.")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn reports_no_content_rather_than_an_empty_string() {
|
||||
let response = transform_response(json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 0}
|
||||
}))
|
||||
.expect("response transforms");
|
||||
assert_eq!(response.choices[0].message.content, None);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn response_carries_no_id_so_python_keeps_its_chatcmpl_id() {
|
||||
let response = transform_response(json!({
|
||||
"id": "msg_should_not_leak",
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [{"type": "text", "text": "hi"}],
|
||||
"stop_reason": "end_turn",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1}
|
||||
}))
|
||||
.expect("response transforms");
|
||||
let value = serde_json::to_value(response).expect("serializable");
|
||||
assert!(
|
||||
value.get("id").is_none(),
|
||||
"the rust response must not carry an id, got {value}"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_a_response_carrying_a_non_text_block() {
|
||||
let err = transform_response(json!({
|
||||
"model": "claude-sonnet-4-5",
|
||||
"content": [{"type": "tool_use", "id": "t1", "name": "f", "input": {}}],
|
||||
"stop_reason": "tool_use",
|
||||
"usage": {"input_tokens": 1, "output_tokens": 1}
|
||||
}))
|
||||
.expect_err("non-text block");
|
||||
assert_eq!(
|
||||
err,
|
||||
CoreError::Unsupported("non-text response content block")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn errors_on_a_response_missing_required_fields() {
|
||||
assert_eq!(
|
||||
transform_response(json!("nope")).expect_err("not an object"),
|
||||
CoreError::InvalidResponse("messages response is not an object".to_string())
|
||||
);
|
||||
assert_eq!(
|
||||
transform_response(json!({"model": "m", "usage": {}})).expect_err("no content"),
|
||||
CoreError::MissingField("content")
|
||||
);
|
||||
assert_eq!(
|
||||
transform_response(json!({"model": "m", "content": []})).expect_err("no usage"),
|
||||
CoreError::MissingField("usage")
|
||||
);
|
||||
assert_eq!(
|
||||
transform_response(json!({"content": [], "usage": {}})).expect_err("no model"),
|
||||
CoreError::MissingField("model")
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_the_messages_url_and_x_api_key_auth() {
|
||||
let config = &ANTHROPIC_CHAT_COMPLETIONS_CONFIG;
|
||||
assert_eq!(
|
||||
config
|
||||
.complete_url(None, "claude-sonnet-4-5", &Map::new(), &|_| None)
|
||||
.expect("url builds"),
|
||||
"https://api.anthropic.com/v1/messages"
|
||||
);
|
||||
assert_eq!(
|
||||
config
|
||||
.auth(Some("sk-x"), "claude-sonnet-4-5", &Map::new(), &|_| None)
|
||||
.expect("auth resolves"),
|
||||
ChatCompletionsAuth::Header {
|
||||
name: "x-api-key",
|
||||
value: "sk-x".to_string()
|
||||
}
|
||||
);
|
||||
assert_eq!(
|
||||
config.default_headers(),
|
||||
&[
|
||||
("anthropic-version", "2023-06-01"),
|
||||
("content-type", "application/json"),
|
||||
]
|
||||
);
|
||||
}
|
||||
|
|
@ -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;
|
||||
|
|
@ -1 +1,2 @@
|
|||
pub mod chat_completions;
|
||||
pub mod messages;
|
||||
|
|
|
|||
|
|
@ -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::*;
|
||||
|
|
|
|||
|
|
@ -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();
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
pub mod transformation;
|
||||
|
|
@ -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), ¶ms(opts))
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_the_converse_body_python_builds() {
|
||||
let body = transform(
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]),
|
||||
json!({"maxTokens": 128, "temperature": 0.2}),
|
||||
);
|
||||
assert_eq!(
|
||||
body,
|
||||
json!({
|
||||
"inferenceConfig": {"maxTokens": 128, "temperature": 0.2},
|
||||
"messages": [{"role": "user", "content": [{"text": "hi"}]}],
|
||||
"system": [{"text": "be terse"}]
|
||||
})
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn always_emits_inference_config_even_when_empty() {
|
||||
let body = transform(json!([{"role": "user", "content": "hi"}]), json!({}));
|
||||
assert_eq!(body["inferenceConfig"], json!({}));
|
||||
assert!(body.get("system").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn places_only_inference_params_in_inference_config() {
|
||||
let body = transform(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({
|
||||
"maxTokens": 64,
|
||||
"temperature": 0.1,
|
||||
"topP": 0.9,
|
||||
"stopSequences": ["STOP"]
|
||||
}),
|
||||
);
|
||||
assert_eq!(
|
||||
body["inferenceConfig"],
|
||||
json!({"maxTokens": 64, "temperature": 0.1, "topP": 0.9, "stopSequences": ["STOP"]})
|
||||
);
|
||||
assert!(body.get("additionalModelRequestFields").is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn merges_consecutive_user_turns_into_one_message() {
|
||||
let body = transform(
|
||||
json!([
|
||||
{"role": "user", "content": "one"},
|
||||
{"role": "user", "content": [{"type": "text", "text": "two"}]},
|
||||
{"role": "assistant", "content": "ack"},
|
||||
{"role": "user", "content": "three"}
|
||||
]),
|
||||
json!({}),
|
||||
);
|
||||
assert_eq!(
|
||||
body["messages"],
|
||||
json!([
|
||||
{"role": "user", "content": [{"text": "one"}, {"text": "two"}]},
|
||||
{"role": "assistant", "content": [{"text": "ack"}]},
|
||||
{"role": "user", "content": [{"text": "three"}]}
|
||||
])
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_streaming() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"stream": true})
|
||||
),
|
||||
Some(Unsupported("streaming"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_top_k_because_python_routes_it_by_base_model() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": "hi"}]),
|
||||
json!({"topK": 40})
|
||||
),
|
||||
Some(Unsupported("unrecognized request parameter"))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_tools_and_other_params_outside_the_allowlist() {
|
||||
for param in [
|
||||
json!({"tools": []}),
|
||||
json!({"tool_choice": {"auto": {}}}),
|
||||
json!({"thinking": {"type": "enabled"}}),
|
||||
json!({"requestMetadata": {"k": "v"}}),
|
||||
json!({"outputConfig": {}}),
|
||||
json!({"_parallel_tool_use_config": {}}),
|
||||
] {
|
||||
assert_eq!(
|
||||
reason(json!([{"role": "user", "content": "hi"}]), param.clone()),
|
||||
Some(Unsupported("unrecognized request parameter")),
|
||||
"expected {param} to decline"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_blank_text_rather_than_substituting_the_anthropic_placeholder() {
|
||||
for content in [
|
||||
json!(""),
|
||||
json!(" "),
|
||||
json!([{"type": "text", "text": " "}]),
|
||||
] {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": content}, {"role": "user", "content": "hi"}]),
|
||||
json!({})
|
||||
),
|
||||
Some(Unsupported("blank message text")),
|
||||
"expected blank content {content} to decline"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_a_message_whose_content_list_is_empty() {
|
||||
// The blank-text check scans parts, so an empty list clears it; Converse
|
||||
// rejects an empty `content` array, which is a decline the core owes the
|
||||
// host before the call rather than an error after it.
|
||||
assert_eq!(
|
||||
reason(json!([{"role": "user", "content": []}]), json!({})),
|
||||
Some(Unsupported("message without content"))
|
||||
);
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([{"role": "user", "content": [{"type": "text", "text": "hi"}]}]),
|
||||
json!({})
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn declines_a_conversation_that_opens_or_closes_on_an_assistant_turn() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([
|
||||
{"role": "assistant", "content": "prefill"},
|
||||
{"role": "user", "content": "hi"}
|
||||
]),
|
||||
json!({})
|
||||
),
|
||||
Some(Unsupported(
|
||||
"conversation does not run user turn to user turn"
|
||||
))
|
||||
);
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "prefill"}
|
||||
]),
|
||||
json!({})
|
||||
),
|
||||
Some(Unsupported(
|
||||
"conversation does not run user turn to user turn"
|
||||
))
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn accepts_a_user_to_user_text_conversation() {
|
||||
assert_eq!(
|
||||
reason(
|
||||
json!([
|
||||
{"role": "system", "content": "be terse"},
|
||||
{"role": "user", "content": "hi"},
|
||||
{"role": "assistant", "content": "hello"},
|
||||
{"role": "user", "content": "again"}
|
||||
]),
|
||||
json!({"maxTokens": 16})
|
||||
),
|
||||
None
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn builds_the_converse_url_from_the_region_in_the_model_id() {
|
||||
let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG;
|
||||
assert_eq!(
|
||||
config
|
||||
.complete_url(None, "us-east-1/anthropic.claude-v2", &Map::new(), &|_| {
|
||||
None
|
||||
})
|
||||
.expect("url builds"),
|
||||
"https://bedrock-runtime.us-east-1.amazonaws.com/model/anthropic.claude-v2/converse"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn falls_back_to_the_region_env_then_the_default_region() {
|
||||
let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG;
|
||||
let with_env = |key: &str| (key == "AWS_REGION_NAME").then(|| "eu-west-1".to_string());
|
||||
assert_eq!(
|
||||
config
|
||||
.complete_url(None, "anthropic.claude-v2", &Map::new(), &with_env)
|
||||
.expect("url builds"),
|
||||
"https://bedrock-runtime.eu-west-1.amazonaws.com/model/anthropic.claude-v2/converse"
|
||||
);
|
||||
assert_eq!(
|
||||
config
|
||||
.complete_url(None, "anthropic.claude-v2", &Map::new(), &|_| None)
|
||||
.expect("url builds"),
|
||||
"https://bedrock-runtime.us-west-2.amazonaws.com/model/anthropic.claude-v2/converse"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn prefers_an_explicit_runtime_endpoint_over_the_api_base() {
|
||||
let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG;
|
||||
let overrides = params(json!({"aws_bedrock_runtime_endpoint": "https://vpce.internal/"}));
|
||||
assert_eq!(
|
||||
config
|
||||
.complete_url(
|
||||
Some("https://ignored.example"),
|
||||
"anthropic.claude-v2",
|
||||
&overrides,
|
||||
&|_| None
|
||||
)
|
||||
.expect("url builds"),
|
||||
"https://vpce.internal/model/anthropic.claude-v2/converse"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn signs_with_sigv4_in_the_resolved_region() {
|
||||
let config = &BEDROCK_CHAT_COMPLETIONS_CONFIG;
|
||||
assert_eq!(
|
||||
config
|
||||
.auth(
|
||||
None,
|
||||
"eu-central-1/anthropic.claude-v2",
|
||||
&Map::new(),
|
||||
&|_| None
|
||||
)
|
||||
.expect("auth resolves"),
|
||||
ChatCompletionsAuth::AwsSigV4 {
|
||||
region: "eu-central-1".to_string()
|
||||
}
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn a_bearer_token_outranks_sigv4_the_way_python_resolves_it() {
|
||||
// Python's get_request_headers reads `api_key` as the Bedrock bearer token
|
||||
// and only falls back to the env when the caller passed none, so each case
|
||||
// pins one of its precedence rules. Signing as the host principal when a
|
||||
// bearer identity is configured would cross an account and quota boundary.
|
||||
let bedrock_env =
|
||||
|key: &str| (key == "AWS_BEARER_TOKEN_BEDROCK").then(|| "from-env".to_string());
|
||||
let no_env = |_: &str| None;
|
||||
let resolve = |api_key, env: &dyn Fn(&str) -> Option<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(¶ms(json!({"aws_access_key_id": "AKIA"}))).is_none());
|
||||
assert!(
|
||||
host_supplied_credentials(¶ms(
|
||||
json!({"aws_access_key_id": " ", "aws_secret_access_key": "s"})
|
||||
))
|
||||
.is_none()
|
||||
);
|
||||
assert!(host_supplied_credentials(&Map::new()).is_none());
|
||||
}
|
||||
|
|
@ -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}", ®ion));
|
||||
let endpoint = endpoint.trim_end_matches('/');
|
||||
// A host that already built the full Converse URL (LiteLLM's Python
|
||||
// path encodes the model id itself) passes it through untouched, the
|
||||
// way the Anthropic config leaves a complete `/v1/messages` URL alone.
|
||||
if endpoint.ends_with(CONVERSE_PATH_SUFFIX) {
|
||||
return Ok(endpoint.to_string());
|
||||
}
|
||||
Ok(format!("{endpoint}/model/{model_id}{CONVERSE_PATH_SUFFIX}"))
|
||||
}
|
||||
|
||||
fn auth(
|
||||
&self,
|
||||
api_key: Option<&str>,
|
||||
model: &str,
|
||||
optional_params: &Map<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;
|
||||
|
|
@ -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";
|
||||
|
|
|
|||
|
|
@ -5,4 +5,5 @@
|
|||
#[cfg(feature = "bedrock-auth")]
|
||||
pub mod audio_transcription;
|
||||
pub mod aws_base;
|
||||
pub mod chat_completions;
|
||||
mod constants;
|
||||
|
|
|
|||
|
|
@ -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(())
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ###########################
|
||||
########################################################################################
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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],
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
100
litellm/integrations/prometheus_metrics_endpoint.py
Normal file
100
litellm/integrations/prometheus_metrics_endpoint.py
Normal 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
|
||||
|
|
@ -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),
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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 (
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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", "")
|
||||
|
|
|
|||
274
litellm/proxy/db/token_auth.py
Normal file
274
litellm/proxy/db/token_auth.py
Normal 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
|
||||
),
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue