Merge remote-tracking branch 'upstream/main' into litellm_lens_live_review

# Conflicts:
#	litellm/proxy/lens/worker.py
#	tests/unit/proxy/lens/test_worker.py
#	ui/litellm-dashboard/src/components/lens/investigations/InvestigationsView.tsx
#	ui/litellm-dashboard/src/components/lens/investigations/detail/InvestigationDetail.tsx
#	ui/litellm-dashboard/src/lib/http/schema.d.ts
This commit is contained in:
Ishaan Jaff 2026-10-03 18:15:36 -07:00
commit ca9db83e17
No known key found for this signature in database
363 changed files with 37139 additions and 15350 deletions

View file

@ -41,6 +41,7 @@ legacy_paths() {
echo tests/unit/enterprise/proxy/hooks
echo tests/unit/enterprise/proxy/management_endpoints
echo tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
echo tests/unit/enterprise/proxy/test_liteadmin.py
echo tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py ;;
enterprise-routing)
echo tests/unit/google_genai

View file

@ -130,37 +130,24 @@ def _unit_selection_arms(repo_root: pathlib.Path = REPO_ROOT) -> Mapping[str, fr
text: Final = _uncommented(script.read_text())
return MappingProxyType(
{
label: frozenset(
match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(body)
)
label: frozenset(match.group(0).rstrip("/") for match in TEST_TOKEN_RE.finditer(body))
for label, body in SELECTION_ARM_RE.findall(text)
}
)
def _unit_selection_tokens(repo_root: pathlib.Path = REPO_ROOT) -> frozenset[str]:
return frozenset(
token for tokens in _unit_selection_arms(repo_root).values() for token in tokens
)
return frozenset(token for tokens in _unit_selection_arms(repo_root).values() for token in tokens)
def _wired_unit_flags(scalars: Iterable[Scalar]) -> frozenset[str]:
return frozenset(
scalar.value
for scalar in scalars
if scalar.key == "unit-flag" and "${{" not in scalar.value
)
return frozenset(scalar.value for scalar in scalars if scalar.key == "unit-flag" and "${{" not in scalar.value)
def _shard_tokens(
scalars: Iterable[Scalar], arms: Mapping[str, frozenset[str]]
) -> frozenset[str]:
def _shard_tokens(scalars: Iterable[Scalar], arms: Mapping[str, frozenset[str]]) -> frozenset[str]:
wired: Final = _wired_unit_flags(scalars)
return _invoked_test_tokens(scalars) | frozenset(
token
for label, tokens in arms.items()
if label in wired
for token in tokens
token for label, tokens in arms.items() if label in wired for token in tokens
)
@ -544,17 +531,37 @@ def _integration_groups(runner: pathlib.Path) -> dict[str, tuple[str, ...]]:
return {group: tuple(folders) for group, folders in ast.literal_eval(mapping).items()}
def _integration_github_files(runner: pathlib.Path) -> frozenset[str]:
module: Final = ast.parse(runner.read_text())
literal: Final = next(
(
node.value
for node in module.body
if isinstance(node, ast.AnnAssign)
and isinstance(node.target, ast.Name)
and node.target.id == "GITHUB_FILES"
),
None,
)
if literal is None:
return frozenset()
values: Final = literal.args[0] if isinstance(literal, ast.Call) else literal
return frozenset(ast.literal_eval(values))
def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozenset[str], tuple[Finding, ...]]:
runner: Final = repo_root / "tests/integration/run.py"
if not runner.exists():
return frozenset(), ()
groups: Final = _integration_groups(runner)
github_files: Final = _integration_github_files(runner)
integration_root: Final = repo_root / "tests/integration"
paths: Final = frozenset(
str(path.relative_to(repo_root))
for folders in groups.values()
for folder in folders
for path in (integration_root / folder).rglob("test_*.py")
if str(path.relative_to(repo_root)) not in github_files
)
browser_manifest: Final = repo_root / "tests/e2e/ui/tests/integrationCritical/expected.json"
browser_nodes: Final = json.loads(browser_manifest.read_text()) if browser_manifest.exists() else ()
@ -595,10 +602,22 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
for path in (repo_root / ".github/workflows").glob("*.y*ml")
for scalar in _scalars(yaml.safe_load(path.read_text()), path.name)
)
findings: Final = tuple(
Finding(path, "integration contract is also selected by GitHub Actions")
for path in paths
if any(_token_covers(token, path) for token in gha_tokens)
findings: Final = (
tuple(
Finding(path, "integration contract is also selected by GitHub Actions")
for path in paths
if any(_token_covers(token, path) for token in gha_tokens)
)
+ tuple(
Finding(path, "GitHub-owned integration contract has no invoking workflow")
for path in sorted(github_files)
if not any(_token_covers(token, path) for token in gha_tokens)
)
+ tuple(
Finding(path, "GitHub-owned integration file is missing")
for path in sorted(github_files)
if not (repo_root / path).is_file()
)
)
browser_commands: Final = tuple(
scalar.value
@ -642,7 +661,7 @@ def _integration_ownership(repo_root: pathlib.Path = REPO_ROOT) -> tuple[frozens
return frozenset(), findings + (
Finding(str(runner.relative_to(repo_root)), "dedicated CircleCI runner is missing"),
)
return paths | browser_paths, findings + group_findings + browser_findings + exclusion_findings
return paths | browser_paths | github_files, findings + group_findings + browser_findings + exclusion_findings
def main() -> int:

View file

@ -15,6 +15,9 @@ on:
- gateway/main.py
- backend/Dockerfile
- backend/main.py
- deploy/lens/**
- litellm/proxy/lens/**
- tests/e2e/migrations/lens_compose_smoke.sh
- docker/component_entrypoint.sh
- docker/entrypoint.sh
- litellm/proxy/prisma_migration.py
@ -37,6 +40,80 @@ concurrency:
cancel-in-progress: true
jobs:
lens-worker-image:
name: lens-worker-image (${{ matrix.arch }})
runs-on: ${{ matrix.runner }}
if: >-
github.event_name != 'pull_request' ||
github.event.pull_request.head.repo.full_name == github.repository
timeout-minutes: 15
permissions:
contents: read
strategy:
fail-fast: false
matrix:
include:
- arch: amd64
runner: ubuntu-latest
grype_sha256: edda0968d8827daab01d32b3cd7de192ae0915005e7bbfcfef9e68e79bc43343
- arch: arm64
runner: ubuntu-24.04-arm
grype_sha256: 553e4c36d9d61349830ba6034d43b8700a7f10576d3e2f4981c0fd2b96086465
steps:
- uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0
with:
persist-credentials: false
- name: Build the release worker
env:
RELEASE_TAG: sha-${{ github.sha }}
run: docker build --build-arg LITELLM_RELEASE_TAG="${RELEASE_TAG}" -f deploy/lens/Dockerfile -t lens-worker-scan .
- name: Verify the standalone worker on a read-only filesystem
env:
RELEASE_TAG: sha-${{ github.sha }}
run: |
docker run --rm --network none --read-only --cap-drop ALL \
--tmpfs /tmp:rw,noexec,nosuid,size=1g --security-opt no-new-privileges \
-e EXPECTED_RELEASE_TAG="${RELEASE_TAG}" --entrypoint python lens-worker-scan -c '
import os
import lens.worker
from lens.release import release_tag
from lens.trace_store import trace_store
assert os.getuid() == 65532
assert release_tag() == os.environ["EXPECTED_RELEASE_TAG"]
with trace_store() as store:
assert store.count() == 0
'
- name: Reject a dependency whose hash has changed
run: |
docker build --target builder -f deploy/lens/Dockerfile -t lens-worker-deps .
sed -E 's/sha256:[0-9a-f]{64}/sha256:0000000000000000000000000000000000000000000000000000000000000000/g' \
deploy/lens/requirements.lock > "$RUNNER_TEMP/tampered.lock"
if docker run --rm -v "$RUNNER_TEMP/tampered.lock:/tmp/tampered.lock:ro" \
--entrypoint uv lens-worker-deps pip sync --python /app/.venv/bin/python \
--require-hashes --only-binary :all: --reinstall --no-cache /tmp/tampered.lock \
> "$RUNNER_TEMP/hash-check.log" 2>&1; then
echo "::error::Dependency hash mismatch was accepted"
exit 1
fi
cat "$RUNNER_TEMP/hash-check.log"
grep -qi 'hash mismatch' "$RUNNER_TEMP/hash-check.log"
- name: Download Grype v0.114.0
env:
ARCH: ${{ matrix.arch }}
GRYPE_SHA256: ${{ matrix.grype_sha256 }}
run: |
curl -fsSL --retry 3 -o "$RUNNER_TEMP/grype.tar.gz" \
"https://github.com/anchore/grype/releases/download/v0.114.0/grype_0.114.0_linux_${ARCH}.tar.gz"
echo "${GRYPE_SHA256} $RUNNER_TEMP/grype.tar.gz" | sha256sum -c -
tar xzf "$RUNNER_TEMP/grype.tar.gz" -C "$RUNNER_TEMP" grype
chmod +x "$RUNNER_TEMP/grype"
- name: Scan the worker for fixable HIGH/CRITICAL CVEs
env:
GRYPE_MATCH_PYTHON_USING_CPES: "true"
run: |
"$RUNNER_TEMP/grype" lens-worker-scan \
--config .grype.yaml --only-fixed --fail-on high --output table
image-scan:
name: image-scan
runs-on: ubuntu-latest
@ -113,7 +190,7 @@ jobs:
persist-credentials: false
- name: Build runtime image
run: docker build -f Dockerfile -t litellm-runtime-scan:${{ github.sha }} .
run: docker build --build-arg LITELLM_RELEASE_TAG=v0.0.0-lens-ci -f Dockerfile -t litellm-runtime-scan:${{ github.sha }} .
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
@ -127,6 +204,11 @@ jobs:
python -m pip install "pytest==9.0.3"
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
- name: Verify the bundled Lens Compose installation and restart
env:
LITELLM_IMAGE: litellm-runtime-scan:${{ github.sha }}
run: bash tests/e2e/migrations/lens_compose_smoke.sh
migrations-image:
name: migrations-image
runs-on: ubuntu-latest

View file

@ -34,7 +34,14 @@ jobs:
with:
persist-credentials: false
- name: Build Lens worker
run: docker build -f deploy/lens/Dockerfile -t lens-worker:${{ github.sha }} .
run: docker build --build-arg LITELLM_RELEASE_TAG=sha-${{ github.sha }} -f deploy/lens/Dockerfile -t lens-worker:${{ github.sha }} .
- name: Reject custom builds without a matching release tag
run: |
if docker build --progress plain -f deploy/lens/Dockerfile -t lens-worker:unversioned . > missing-tag.log 2>&1; then
echo "::error::An unversioned worker build unexpectedly succeeded"
exit 1
fi
grep -F 'LITELLM_RELEASE_TAG: Pass --build-arg LITELLM_RELEASE_TAG matching the gateway' missing-tag.log
- name: Verify standalone imports with a read-only filesystem
run: |
docker run --rm --network none --read-only --cap-drop ALL --tmpfs /tmp:rw,noexec,nosuid,size=1g \
@ -54,11 +61,11 @@ jobs:
-v "$PWD/tests/proxy_behavior/lens/worker_storage_smoke.py:/app/storage_smoke.py:ro" \
--entrypoint python lens-worker:${{ github.sha }} /app/storage_smoke.py
- name: Publish versioned Lens worker
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm'
if: github.event_name != 'pull_request' && github.repository == 'BerriAI/litellm' && github.ref == 'refs/heads/main'
env:
REGISTRY_TOKEN: ${{ secrets.GITHUB_TOKEN }}
REGISTRY_USER: ${{ github.actor }}
IMAGE: ghcr.io/berriai/litellm-lens-worker:sha-${{ github.sha }}
IMAGE: ghcr.io/berriai/litellm-lens-worker-dev:sha-${{ github.sha }}
run: |
printf '%s' "$REGISTRY_TOKEN" | docker login ghcr.io -u "$REGISTRY_USER" --password-stdin
docker tag lens-worker:${{ github.sha }} "$IMAGE"

View file

@ -45,6 +45,13 @@ jobs:
fail-fast: false
matrix:
include:
- shard: roi-database
test-path: "tests/integration/database/test_roi_observed.py"
seed: none
workers: 0
timeout-minutes: 10
job-timeout-minutes: 35
- shard: proxy-behavior
test-path: "tests/proxy_behavior"
seed: db-push
@ -147,7 +154,7 @@ jobs:
env:
TEST_PATH: ${{ matrix.test-path }}
WORKERS: ${{ matrix.workers }}
PYTEST_ADDOPTS: ${{ matrix.shard == 'proxy-behavior' && '--cov=./litellm --cov-report=xml:coverage-lens-postgres.xml' || '' }}
PYTEST_ADDOPTS: ${{ matrix.shard == 'proxy-behavior' && '--cov=./litellm --cov-report=xml:coverage-lens-postgres.xml' || matrix.shard == 'roi-database' && '--cov=./litellm --cov-report=xml:coverage-roi-postgres.xml' || '' }}
run: |
if [ "${WORKERS}" = "0" ]; then
uv run --no-sync pytest ${TEST_PATH:?} -vv --tb=short --durations=10
@ -165,3 +172,14 @@ jobs:
files: coverage-lens-postgres.xml
flags: lens-postgres
fail_ci_if_error: true
- name: Upload ROI database coverage
if: steps.changes.outputs.decision != 'skip' && matrix.shard == 'roi-database' && !cancelled()
uses: codecov/codecov-action@303a32d7a59b442fa8d48b6a1cc6825c09c847a5 # v7.1.1
with:
use_oidc: true
version: v11.3.1
root_dir: ${{ github.workspace }}
files: coverage-roi-postgres.xml
flags: roi-postgres
fail_ci_if_error: true

View file

@ -114,8 +114,20 @@ RUN HOME=/opt/prisma XDG_CACHE_HOME=/opt/prisma/.cache PRISMA_BINARY_CACHE_DIR=/
RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \
sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
FROM $LITELLM_BUILD_IMAGE AS liteadmin-builder
COPY --from=uvbin /uv /usr/local/bin/uv
RUN apk add --no-cache python-3.13
ADD --checksum=sha256:2f7ae5cdd9d91731c0990e74a58239dc3e3fd2bf28dab23b55eafcdc47aaf87e \
https://github.com/BerriAI/litellm-admin-agent/archive/ef501e94bc9fbacb9233b922abf71427f030408c.tar.gz /tmp/liteadmin.tar.gz
RUN mkdir /tmp/liteadmin && tar xzf /tmp/liteadmin.tar.gz --strip-components=1 -C /tmp/liteadmin && \
uv venv /opt/liteadmin --python python3.13 && \
uv pip install --python /opt/liteadmin/bin/python --require-hashes -r /tmp/liteadmin/requirements.txt && \
uv pip install --python /opt/liteadmin/bin/python --no-deps /tmp/liteadmin
# Runtime stage
FROM $LITELLM_RUNTIME_IMAGE AS runtime
ARG LITELLM_RELEASE_TAG=""
ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG}
USER root
@ -141,6 +153,7 @@ ENV PATH="/app/.venv/bin:${PATH}" \
# ship (manifest-scanning tools attribute everything in it to this image).
# entrypoint.sh invokes litellm/proxy/prisma_migration.py by source path.
COPY --from=builder /app/.venv /app/.venv
COPY --from=liteadmin-builder /opt/liteadmin /opt/liteadmin
COPY --from=builder /app/docker /app/docker
COPY --from=builder /app/schema.prisma /app/schema.prisma
COPY --from=builder /app/litellm/proxy/prisma_migration.py /app/litellm/proxy/prisma_migration.py

View file

@ -58,7 +58,7 @@ help:
@echo " make test-unit-helm - Run helm unit tests"
@echo " make test-rust-extension - Build the Rust extension and run its public Python tests"
@echo " make rust-sqlx-prepare - Refresh litellm-rust/crates/db/.sqlx against a migrated Postgres container"
@echo " make lens-dev - Run proxy + Lens worker + hot-reload dashboard (LENS_DEV_PROXY_PORT, LENS_DEV_UI_PORT)"
@echo " make lens-dev - Run proxy + Lens worker + hot-reload dashboard (ARGS=\"--seed large\", LENS_DEV_PROXY_PORT, LENS_DEV_UI_PORT)"
@echo ""
@echo "Heavy targets (check, lint) queue for LITELLM_GATE_SLOTS machine-wide"
@echo "slots (default 2; 0 disables) so parallel sessions don't thrash one machine."
@ -313,7 +313,7 @@ rust-sqlx-prepare:
cd litellm-rust && cargo run -p litellm-db-testing --bin sqlx-prepare
lens-dev:
./scripts/lens_dev.sh
./scripts/lens_dev.sh $(ARGS)
test: install-test-deps
$(UV_RUN) pytest tests/

View file

@ -71,6 +71,8 @@ RUN sed -i 's/\r$//' docker/component_entrypoint.sh && chmod +x docker/component
# ---------- Runtime ----------
FROM $LITELLM_RUNTIME_IMAGE AS runtime
ARG LITELLM_RELEASE_TAG=""
ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG}
USER root

View file

@ -22,6 +22,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
"/customer/",
"/end_user/",
"/sso/",
"/liteadmin/slack/connect/",
"/login",
"/v2/login",
"/v3/login",

View file

@ -5710,6 +5710,17 @@
}
},
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "histogram_quantile(0.95, sum(rate(litellm_anthropic_wif_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "anthropic_wif",
"range": true,
"refId": "A"
},
{
"datasource": {
"type": "prometheus",
@ -5719,7 +5730,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_auth_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "auth",
"range": true,
"refId": "A"
"refId": "B"
},
{
"datasource": {
@ -5730,7 +5741,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_batch_write_to_db_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "batch_write_to_db",
"range": true,
"refId": "B"
"refId": "C"
},
{
"datasource": {
@ -5741,7 +5752,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_postgres_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "postgres",
"range": true,
"refId": "C"
"refId": "D"
},
{
"datasource": {
@ -5752,7 +5763,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_proxy_pre_call_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "proxy_pre_call",
"range": true,
"refId": "D"
"refId": "E"
},
{
"datasource": {
@ -5763,7 +5774,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_redis_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "redis",
"range": true,
"refId": "E"
"refId": "F"
},
{
"datasource": {
@ -5774,7 +5785,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_redis_daily_org_spend_update_queue_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "redis_daily_org_spend_update_queue",
"range": true,
"refId": "F"
"refId": "G"
},
{
"datasource": {
@ -5785,7 +5796,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_redis_daily_tag_spend_update_queue_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "redis_daily_tag_spend_update_queue",
"range": true,
"refId": "G"
"refId": "H"
},
{
"datasource": {
@ -5796,7 +5807,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_redis_daily_team_spend_update_queue_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "redis_daily_team_spend_update_queue",
"range": true,
"refId": "H"
"refId": "I"
},
{
"datasource": {
@ -5807,7 +5818,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_redis_window_spend_update_queue_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "redis_window_spend_update_queue",
"range": true,
"refId": "I"
"refId": "J"
},
{
"datasource": {
@ -5818,7 +5829,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_reset_budget_job_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "reset_budget_job",
"range": true,
"refId": "J"
"refId": "K"
},
{
"datasource": {
@ -5829,7 +5840,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_router_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "router",
"range": true,
"refId": "K"
"refId": "L"
},
{
"datasource": {
@ -5840,7 +5851,7 @@
"expr": "histogram_quantile(0.95, sum(rate(litellm_self_latency_bucket[$__rate_interval])) by (le))",
"legendFormat": "self",
"range": true,
"refId": "L"
"refId": "M"
}
],
"title": "Service latency p95 (litellm_<service>_latency)",
@ -5888,6 +5899,28 @@
}
},
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "sum(rate(litellm_anthropic_wif_total_requests_total[$__rate_interval]))",
"legendFormat": "anthropic_wif",
"range": true,
"refId": "A"
},
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "sum(rate(litellm_anthropic_wif_cache_total_requests_total[$__rate_interval]))",
"legendFormat": "anthropic_wif_cache",
"range": true,
"refId": "B"
},
{
"datasource": {
"type": "prometheus",
@ -5897,7 +5930,7 @@
"expr": "sum(rate(litellm_auth_total_requests_total[$__rate_interval]))",
"legendFormat": "auth",
"range": true,
"refId": "A"
"refId": "C"
},
{
"datasource": {
@ -5908,7 +5941,7 @@
"expr": "sum(rate(litellm_batch_write_to_db_total_requests_total[$__rate_interval]))",
"legendFormat": "batch_write_to_db",
"range": true,
"refId": "B"
"refId": "D"
},
{
"datasource": {
@ -5919,7 +5952,7 @@
"expr": "sum(rate(litellm_postgres_total_requests_total[$__rate_interval]))",
"legendFormat": "postgres",
"range": true,
"refId": "C"
"refId": "E"
},
{
"datasource": {
@ -5930,7 +5963,7 @@
"expr": "sum(rate(litellm_proxy_pre_call_total_requests_total[$__rate_interval]))",
"legendFormat": "proxy_pre_call",
"range": true,
"refId": "D"
"refId": "F"
},
{
"datasource": {
@ -5941,7 +5974,7 @@
"expr": "sum(rate(litellm_redis_total_requests_total[$__rate_interval]))",
"legendFormat": "redis",
"range": true,
"refId": "E"
"refId": "G"
},
{
"datasource": {
@ -5952,7 +5985,7 @@
"expr": "sum(rate(litellm_redis_daily_org_spend_update_queue_total_requests_total[$__rate_interval]))",
"legendFormat": "redis_daily_org_spend_update_queue",
"range": true,
"refId": "F"
"refId": "H"
},
{
"datasource": {
@ -5963,7 +5996,7 @@
"expr": "sum(rate(litellm_redis_daily_tag_spend_update_queue_total_requests_total[$__rate_interval]))",
"legendFormat": "redis_daily_tag_spend_update_queue",
"range": true,
"refId": "G"
"refId": "I"
},
{
"datasource": {
@ -5974,7 +6007,7 @@
"expr": "sum(rate(litellm_redis_daily_team_spend_update_queue_total_requests_total[$__rate_interval]))",
"legendFormat": "redis_daily_team_spend_update_queue",
"range": true,
"refId": "H"
"refId": "J"
},
{
"datasource": {
@ -5985,7 +6018,7 @@
"expr": "sum(rate(litellm_redis_window_spend_update_queue_total_requests_total[$__rate_interval]))",
"legendFormat": "redis_window_spend_update_queue",
"range": true,
"refId": "I"
"refId": "K"
},
{
"datasource": {
@ -5996,7 +6029,7 @@
"expr": "sum(rate(litellm_reset_budget_job_total_requests_total[$__rate_interval]))",
"legendFormat": "reset_budget_job",
"range": true,
"refId": "J"
"refId": "L"
},
{
"datasource": {
@ -6007,7 +6040,7 @@
"expr": "sum(rate(litellm_router_total_requests_total[$__rate_interval]))",
"legendFormat": "router",
"range": true,
"refId": "K"
"refId": "M"
},
{
"datasource": {
@ -6018,7 +6051,7 @@
"expr": "sum(rate(litellm_self_total_requests_total[$__rate_interval]))",
"legendFormat": "self",
"range": true,
"refId": "L"
"refId": "N"
}
],
"title": "Service request rate (litellm_<service>_total_requests)",
@ -6066,6 +6099,28 @@
}
},
"targets": [
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "sum(rate(litellm_anthropic_wif_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "anthropic_wif / {{error_class}}",
"range": true,
"refId": "A"
},
{
"datasource": {
"type": "prometheus",
"uid": "${DS_PROMETHEUS}"
},
"editorMode": "code",
"expr": "sum(rate(litellm_anthropic_wif_cache_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "anthropic_wif_cache / {{error_class}}",
"range": true,
"refId": "B"
},
{
"datasource": {
"type": "prometheus",
@ -6075,7 +6130,7 @@
"expr": "sum(rate(litellm_auth_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "auth / {{error_class}}",
"range": true,
"refId": "A"
"refId": "C"
},
{
"datasource": {
@ -6086,7 +6141,7 @@
"expr": "sum(rate(litellm_batch_write_to_db_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "batch_write_to_db / {{error_class}}",
"range": true,
"refId": "B"
"refId": "D"
},
{
"datasource": {
@ -6097,7 +6152,7 @@
"expr": "sum(rate(litellm_postgres_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "postgres / {{error_class}}",
"range": true,
"refId": "C"
"refId": "E"
},
{
"datasource": {
@ -6108,7 +6163,7 @@
"expr": "sum(rate(litellm_proxy_pre_call_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "proxy_pre_call / {{error_class}}",
"range": true,
"refId": "D"
"refId": "F"
},
{
"datasource": {
@ -6119,7 +6174,7 @@
"expr": "sum(rate(litellm_redis_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "redis / {{error_class}}",
"range": true,
"refId": "E"
"refId": "G"
},
{
"datasource": {
@ -6130,7 +6185,7 @@
"expr": "sum(rate(litellm_redis_daily_org_spend_update_queue_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "redis_daily_org_spend_update_queue / {{error_class}}",
"range": true,
"refId": "F"
"refId": "H"
},
{
"datasource": {
@ -6141,7 +6196,7 @@
"expr": "sum(rate(litellm_redis_daily_tag_spend_update_queue_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "redis_daily_tag_spend_update_queue / {{error_class}}",
"range": true,
"refId": "G"
"refId": "I"
},
{
"datasource": {
@ -6152,7 +6207,7 @@
"expr": "sum(rate(litellm_redis_daily_team_spend_update_queue_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "redis_daily_team_spend_update_queue / {{error_class}}",
"range": true,
"refId": "H"
"refId": "J"
},
{
"datasource": {
@ -6163,7 +6218,7 @@
"expr": "sum(rate(litellm_redis_window_spend_update_queue_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "redis_window_spend_update_queue / {{error_class}}",
"range": true,
"refId": "I"
"refId": "K"
},
{
"datasource": {
@ -6174,7 +6229,7 @@
"expr": "sum(rate(litellm_reset_budget_job_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "reset_budget_job / {{error_class}}",
"range": true,
"refId": "J"
"refId": "L"
},
{
"datasource": {
@ -6185,7 +6240,7 @@
"expr": "sum(rate(litellm_router_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "router / {{error_class}}",
"range": true,
"refId": "K"
"refId": "M"
},
{
"datasource": {
@ -6196,7 +6251,7 @@
"expr": "sum(rate(litellm_self_failed_requests_total[$__rate_interval])) by (error_class)",
"legendFormat": "self / {{error_class}}",
"range": true,
"refId": "L"
"refId": "N"
}
],
"title": "Service failure rate (litellm_<service>_failed_requests)",

View file

@ -1,6 +1,6 @@
# LiteLLM All Prometheus Metrics dashboard
Every `litellm_*` metric family the proxy can expose on `/metrics` (136 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
Every `litellm_*` metric family the proxy can expose on `/metrics` (141 families across 97 panels), grouped into rows: proxy traffic, latency, spend and tokens, cache, LLM API deployments, key and team rate limits, budgets, guardrails, MCP, managed files and batches, users and teams, the Redis circuit breaker, the spend log cleanup job, and the `prometheus_system` service callback metrics (per-service latency, request and failure rates, spend update queue sizes). Panel titles are the metric names so you can grep the JSON for the metric you care about
Import `grafana_dashboard.json` from **Dashboards > New > Import** and pick your Prometheus data source when prompted (the `DS_PROMETHEUS` variable). Counters are plotted as `rate()` over `$__rate_interval`, histograms as p50 / p95 / p99, gauges as the raw value grouped by the most useful label. Every query names the metric exactly as the proxy emits it (counters carry the `_total` suffix the Prometheus client adds), and `tests/unit/integrations/test_prometheus_metric_name_consistency.py` fails if a metric is renamed without updating this dashboard

View file

@ -1,7 +1,28 @@
FROM python:3.12-slim
ARG LITELLM_BUILD_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG LITELLM_RUNTIME_IMAGE=cgr.dev/chainguard/wolfi-base@sha256:1d95114038f76513a9ace6fca107d5582b08c65981f81f61cb56bf7fd2ef216d
ARG UV_IMAGE=ghcr.io/astral-sh/uv:0.11.7@sha256:240fb85ab0f263ef12f492d8476aa3a2e4e1e333f7d67fbdd923d00a506a516a
FROM $UV_IMAGE AS uvbin
FROM $LITELLM_BUILD_IMAGE AS builder
COPY --from=uvbin /uv /usr/local/bin/uv
RUN apk add --no-cache python-3.13
ENV UV_PYTHON_DOWNLOADS=0 UV_LINK_MODE=copy
WORKDIR /app
RUN pip install --no-cache-dir httpx==0.28.1 pydantic==2.11.7
COPY litellm/proxy/lens/__init__.py litellm/proxy/lens/models.py litellm/proxy/lens/trace_store.py litellm/proxy/lens/analysis.py litellm/proxy/lens/worker.py /app/lens/
COPY deploy/lens/requirements.lock /tmp/requirements.lock
RUN uv venv --python python3.13 /app/.venv && \
uv pip sync --python /app/.venv/bin/python --require-hashes --only-binary :all: /tmp/requirements.lock
FROM $LITELLM_RUNTIME_IMAGE AS runtime
ARG LITELLM_RELEASE_TAG=""
RUN : "${LITELLM_RELEASE_TAG:?Pass --build-arg LITELLM_RELEASE_TAG matching the gateway}"
RUN apk add --no-cache python-3.13
ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG} \
PATH="/app/.venv/bin:${PATH}" \
PYTHONDONTWRITEBYTECODE=1
WORKDIR /app
COPY --from=builder /app/.venv /app/.venv
COPY litellm/proxy/lens/__init__.py litellm/proxy/lens/models.py litellm/proxy/lens/trace_store.py litellm/proxy/lens/analysis.py litellm/proxy/lens/worker.py litellm/proxy/lens/release.py /app/lens/
COPY litellm/proxy/lens/prompts/ /app/lens/prompts/
USER 65532:65532
CMD ["python", "-m", "lens.worker"]

View file

@ -1,4 +1,7 @@
**
!deploy/
!deploy/lens/
!deploy/lens/requirements.lock
!litellm/
!litellm/proxy/
!litellm/proxy/lens/
@ -7,5 +10,6 @@
!litellm/proxy/lens/trace_store.py
!litellm/proxy/lens/analysis.py
!litellm/proxy/lens/worker.py
!litellm/proxy/lens/release.py
!litellm/proxy/lens/prompts/
!litellm/proxy/lens/prompts/**

View file

@ -2,7 +2,60 @@
Lens reviews recorded activity and saves evidence-linked findings in the LiteLLM dashboard under Observability, Lens (`/ui/lens/`)
## Start a worker
## Install the release stack
Each stable, RC, and dev release containing Lens publishes the worker at the same version on GHCR and Docker Hub. Use the [LiteLLM releases page](https://github.com/BerriAI/litellm/releases) to select a version that includes the coordinated worker release
For a new local installation, install Docker with Compose, download the two release files, and create a private environment file. Replace `X.Y.Z` with the release version, without `v` (RCs use `X.Y.Z-rc.N`)
```bash
mkdir litellm-lens
cd litellm-lens
LENS_RELEASE=X.Y.Z
curl -fSLo compose.yaml "https://raw.githubusercontent.com/BerriAI/litellm/v${LENS_RELEASE}/deploy/lens/stack.yaml"
curl -fSLo config.yaml "https://raw.githubusercontent.com/BerriAI/litellm/v${LENS_RELEASE}/deploy/lens/config.yaml"
umask 077
printf 'LITELLM_VERSION=%s\nLITELLM_MASTER_KEY=sk-%s\nLITELLM_SALT_KEY=sk-%s\n' \
"$LENS_RELEASE" "$(openssl rand -hex 32)" "$(openssl rand -hex 32)" > .env
printf 'POSTGRES_PASSWORD=%s\nCLICKHOUSE_PASSWORD=%s\n' \
"$(openssl rand -hex 32)" "$(openssl rand -hex 32)" >> .env
docker compose up -d
```
Open `http://localhost:4000/ui/`, log in as `admin` with `LITELLM_MASTER_KEY` from `.env`, and add a model in the dashboard. In Lens, select **Connect worker**, choose that model and a monthly budget, then **Get install command**. Expand **Using Docker Compose or Helm?**, copy the worker token, and add `LENS_WORKER_TOKEN=<token>` to `.env`
```bash
docker compose --profile lens up -d
```
The stack starts LiteLLM, PostgreSQL, ClickHouse, and the worker from published images. The dashboard shows **Worker connected**. The worker has a limited token, no database credentials, and no provider keys. The stack exposes only the dashboard on localhost; use your normal ingress and managed databases for a public production deployment
Keep `.env` private and preserve its salt key. Keep both named database volumes. To upgrade, wait for active investigations to finish, stop the worker, change only `LITELLM_VERSION`, then pull and recreate the stack:
```bash
docker compose --profile lens stop lens-worker
# Update LITELLM_VERSION in .env to the new release
docker compose --profile lens pull
docker compose --profile lens up -d
```
This preserves your investigations, findings, model credentials, and worker token. Never use `down -v` during an upgrade. If moving from an existing installation, keep its databases and add the standalone worker instead of creating an empty replacement stack
## Helm
The componentized `helm/litellm` chart includes an optional Lens worker. Configure PostgreSQL and ClickHouse as usual, install the chart, then obtain a limited worker token from Lens setup. Store it in a Kubernetes Secret and enable the worker in your values:
```yaml
lensWorker:
enabled: true
tokenSecret:
name: litellm-lens-worker
key: token
```
Published release charts pin the worker's approved image digest. Source charts without a digest default to the chart's application version. The chart connects the worker to the backend service. Keep these values and the Secret when upgrading the chart so the gateway and worker upgrade together. `lensWorker.replicaCount` controls simultaneous investigations. To use a private registry or external proxy, set `lensWorker.image.repository`, `lensWorker.image.digest` (or `tag` for a source build), and `lensWorker.url`. A digest takes precedence over the tag. The dashboard uses the chart's worker image for standalone install commands too
## Standalone worker
Upgrade your existing LiteLLM proxy to a release that includes Lens with PostgreSQL and agent tracing. Configure one ClickHouse URL for trace writes, bounded reads, and Lens queries:
@ -23,17 +76,17 @@ In **Lens > Investigations**, click **Connect worker**, choose an analysis model
The command already contains the compatible worker image and one worker token. The selected virtual key stays on the proxy; its secret is never sent to the worker. No source checkout, environment file, or second LiteLLM deployment is needed. Keep the command private because it includes the token. The LiteLLM release provides the dashboard and APIs; the container only runs background analysis
The dashboard and Compose file pin a verified worker image by digest. The image uses Linux amd64, and the generated command selects that platform. CI also publishes immutable `:sha-<commit>` tags for successful worker builds on `main`. Keep the worker image compatible with your gateway version
The dashboard selects the worker image matching the running gateway release. Release images support Linux amd64 and arm64. CI also publishes `:sha-<commit>` development images; use those only with a gateway built from the same commit and release tag
After upgrading the gateway, update the worker image and redeploy it while keeping its proxy URL and token. Existing containers do not update automatically. If an investigation reports a worker compatibility error, update the image before retrying
For deployments managed with Compose, download `compose.yaml` and provide `LITELLM_URL` and `LENS_WORKER_TOKEN` in an environment file. Its default image is already selected:
For deployments managed with Compose, download `compose.yaml` and provide `LITELLM_URL`, `LENS_WORKER_TOKEN`, and `LITELLM_VERSION` (without `v`) in a private environment file. To use another registry, set `LENS_WORKER_IMAGE` to the compatible image instead of setting a version:
```bash
docker compose --env-file /path/to/lens.env -f compose.yaml up -d
```
Developers can build locally with `LENS_WORKER_IMAGE=litellm-lens-worker:local docker compose -f deploy/lens/compose.yaml -f deploy/lens/compose.build.yaml up -d --build`. To work on Lens itself, `make lens-dev` runs the proxy, a worker from source and the hot-reload dashboard together; set `LENS_DEV_PROXY_PORT` / `LENS_DEV_UI_PORT` to move them off 4000/3000
To work on Lens itself, `make lens-dev` runs the proxy, a worker from source and the hot-reload dashboard together; set `LENS_DEV_PROXY_PORT` / `LENS_DEV_UI_PORT` to move them off 4000/3000. For a local container build, set `LENS_WORKER_IMAGE=litellm-lens-worker:local` and `LITELLM_RELEASE_TAG` to the gateway's release tag, then use `docker compose -f deploy/lens/compose.yaml -f deploy/lens/compose.build.yaml up -d --build`
The generated command gives the worker 1 GiB of temporary memory-backed storage, shared across parallel reviews. Change `size=1g` in the Docker command or set `LENS_WORKER_TMP_SIZE` with Compose to fit your server and workload. A storage failure marks the scan as failed, cleans up temporary traces, and leaves the worker available for other scans; it does not silently truncate the review. Existing workers must be recreated with the new image and mount options
@ -107,6 +160,38 @@ curl "$LITELLM_URL/lens/$LENS_ID/runs/$BATCH_ID" -H "Authorization: Bearer $LITE
Creation queues the first batch. Posting to `/lens/{id}/runs` queues another, or returns the existing active batch. The run response contains its ID under `jobs[0].id`. Poll the batch URL for status, findings and assessments. List responses omit large result payloads; request a batch to retrieve them. Supply an optional complete `settings` object on the runs POST for a one-off override; the saved lens stays unchanged. Selection accepts `team_id`, exact `filters`, and opaque `execution_ids` returned by `/lens/preview/sample`. Preview accepts `offset` and `as_of` to keep the time window fixed while paging. Feedback uses `PATCH /lens/{id}/findings/{finding_id}` with `status` and `reason`
## Local development
`make lens-dev ARGS=--seed` starts the full dev stack. The live dashboard is at `http://localhost:3000/ui/lens/`, with login at `http://localhost:3000/ui/login/`. Next.js forwards API requests to the proxy on port 4000, so login and navigation stay in the live UI and edits hot-reload
The default is Next.js dev with no production build (`LENS_DEV_BUILD_UI=0`). Set `LENS_DEV_BUILD_UI=1` when you also want a fresh static dashboard at `http://localhost:4000/ui/`. Build output goes to `.lens-dev/logs/ui-build.log`; a failed build stops startup. Both modes keep the live dashboard on port 3000. Startup checks the live login route before seeding and fails with the UI log path if Next.js exits. `LENS_DEV_STARTUP_TIMEOUT_SECONDS` controls startup readiness retries (default 300; `LENS_DEV_READINESS_REQUEST_TIMEOUT_SECONDS` caps each HTTP probe, default 5)
For local fixture data, run `make lens-dev ARGS=--seed`. Use `make lens-dev ARGS="--seed large"` for 2,000 fixture copies, over one million spans and linked request logs. To seed a running stack without restarting it, use `make lens-dev ARGS="--seed-only --seed large --copies 100"`. The default profile replays one copy of every checked-in capture through authenticated `/v1/traces`, including failures, retries, streaming and multiple agent frameworks. Large seeds use the same parser and compressed ClickHouse writer in batches of four copies, and write matching request logs to PostgreSQL. The first and last batches verify linked spend totals through the proxy
Seeds append fresh IDs on every invocation and spread copies over recent timestamps. Restarts without `SEED` do not add data. Lens excludes activity received in the last two minutes, so wait two minutes after seeding before checking investigation previews. `LENS_DEV_SEED_COPIES` overrides total copies, and `LENS_DEV_SEED_BATCH_COPIES` overrides copies per bulk insert (default 4, about 2,000 spans). Start with four or fewer on a constrained machine. Larger batches still respect the existing ClickHouse insert size limit; each capture is decoded separately within the OTLP safety budget. Large seeds test data volume and pagination, rather than concurrent ingestion throughput or review accuracy. They can use substantial disk space; adjust `--copies` for your machine. Seeding expects the generated local tracing configuration. The old `run_tracing_proxy_local.sh --seed` command forwards to Lens dev, using its ports and saved master key
Local ingestion limits are explicit and configurable. Set OTLP and ClickHouse variables before starting the proxy and seeder so both processes use the same settings. Invalid, zero and negative values fail instead of silently falling back. Changing these limits does not require rebuilding Rust
| Environment variable | Default | Controls |
| --- | --- | --- |
| `LENS_DEV_SEED_COPIES` | 1 default, 2000 large | Total fixture copies |
| `LENS_DEV_SEED_BATCH_COPIES` | 4 | Copies per bulk insert |
| `LENS_DEV_SEED_TIMEOUT_SECONDS` | 120 | Seeder HTTP timeout |
| `OTLP_MAX_BODY_BYTES` | 16777216 | HTTP body and decompressed payload bytes |
| `OTLP_MAX_CONCURRENT_INGESTS` | 2 | Concurrent proxy ingestion requests |
| `OTLP_MAX_ATTRIBUTE_VALUE_BYTES` | 65536 | Stored attribute/content bytes |
| `OTLP_MAX_DECODE_DEPTH` | 32 | Nested decode depth |
| `OTLP_MAX_DECODE_NODES` | 65536 | JSON values or protobuf fields per export |
| `OTLP_MAX_SPANS` | 4096 | Spans per export |
| `OTLP_MAX_ATTRIBUTES` | 256 | Attributes per resource, scope, span, event or link |
| `OTLP_MAX_EVENTS` | 256 | Events per span |
| `OTLP_MAX_LINKS` | 256 | Links per span |
| `OTLP_MAX_DECODED_SPAN_BYTES` | 16777216 | Decoded span allocation budget |
| `CLICKHOUSE_TRACE_MAX_INSERT_BYTES` | 67108864 | Encoded trace or spend insert bytes |
| `CLICKHOUSE_INSERT_TIMEOUT_SECONDS` | 30 | ClickHouse insert HTTP timeout |
The wire parsers also enforce their library recursion limits (128 levels for JSON, 100 for protobuf). Raising the configured depth does not remove those parser limits. Bulk seeding parses each capture separately, keeping the per-export limits distinct from the bulk insert limit. Use smaller batches if an insert exceeds its byte budget. For example, `LENS_DEV_SEED_COPIES=100 LENS_DEV_SEED_BATCH_COPIES=2 make lens-dev ARGS="--seed large"`
## Quality evaluation
Run the checked-in cases against a configured real model. Expected labels are used only for scoring, never passed to the model. Dev and held-out cases include missing outcomes, failed tools, recovery, handoffs, unsupported claims, repeated work, long evidence and prompt injection. The background option adds clean arithmetic traces to test rare-issue discovery at scale; those repeated synthetic cases do not establish accuracy on every production workload
@ -130,3 +215,19 @@ The Lens API now uses `/lens` instead of `/engine`, list responses use `lenses`,
Stop workers and let active scans finish before upgrading. Deploy proxy instances together: older proxies cannot use the renamed database tables. The schema migration renames the three Lens tables and the run-history identifier column in place, preserving saved investigations, findings, history, worker credentials, and billing assignments. Existing migration files retain their original names and checksums
Upgrades using `--use_prisma_db_push` stop before schema changes if any legacy Lens table exists, preventing Prisma from dropping saved data. Apply `litellm-proxy-extras/litellm_proxy_extras/migrations/20261001100000_rename_lens/migration.sql` to the configured database schema before retrying. Deployments already using migration history can instead start without `--use_prisma_db_push` to apply the shipped migration normally. Fresh databases and databases already using the renamed tables can continue using database push
## Release compatibility
Released gateway and worker images carry `LITELLM_RELEASE_TAG`. A worker announces its release and protocol before claiming an investigation. A mismatch returns HTTP 409 with the required image, leaving queued investigations untouched. During a rolling upgrade, workers wait for a gateway from their release
The dashboard reads its image from the running gateway. `LENS_WORKER_IMAGE` overrides the registry/image for private deployments. Worker-only Compose accepts `LITELLM_VERSION` (without `v`) or an explicit `LENS_WORKER_IMAGE`. Release workers are available as `ghcr.io/berriai/litellm-lens-worker:vX.Y.Z` and `docker.io/litellm/litellm-lens-worker:vX.Y.Z`, including matching RC/dev suffixes, on amd64 and arm64
For source development, use `make lens-dev`, which gives the proxy and source worker the same commit identity. For custom containers, build both from the same checkout with `--build-arg LITELLM_RELEASE_TAG=sha-$(git rev-parse HEAD)` and set the proxy's `LENS_WORKER_IMAGE` to the worker image you built. An unlabelled custom build refuses worker setup and claims instead of guessing from the Python package version. Normal package-index installations use their installed release version
The hourly development pipeline pins all component images to the same selected commit and publishes its chart only after every build and worker smoke test succeeds. The public commit-tagged worker workflow publishes to `ghcr.io/berriai/litellm-lens-worker-dev` on Lens-related changes, so an arbitrary `main` commit may require building your own pair; do not substitute the newest available worker
## Worker dependencies
The worker uses the same digest-pinned Wolfi base and Python version as the component images. Python dependencies and their hashes are locked in `deploy/lens/requirements.lock`. To update them, edit `deploy/lens/requirements.in`, then run `uv pip compile --universal --python-version 3.13 --generate-hashes --no-emit-index-url deploy/lens/requirements.in -o deploy/lens/requirements.lock`. The image installs only the locked wheels with hash verification. CI builds and scans both native architectures

View file

@ -3,4 +3,6 @@ services:
build:
context: ../..
dockerfile: deploy/lens/Dockerfile
args:
LITELLM_RELEASE_TAG: ${LITELLM_RELEASE_TAG:?Set the release tag used by the gateway}
image: litellm-lens-worker:local

View file

@ -1,6 +1,6 @@
services:
lens-worker:
image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker@sha256:44f0597c7583dcfef999ece9a8bc02cfeb9f0f5167a1221cee3bd10b1b79271b}
image: ${LENS_WORKER_IMAGE:-ghcr.io/berriai/litellm-lens-worker:v${LITELLM_VERSION:?Set LITELLM_VERSION to the gateway release, without the v prefix}}
environment:
LITELLM_URL: ${LITELLM_URL:?Set the URL reachable from this container}
LENS_WORKER_TOKEN: ${LENS_WORKER_TOKEN:?Create a worker credential in the Lens UI}

7
deploy/lens/config.yaml Normal file
View file

@ -0,0 +1,7 @@
general_settings:
master_key: os.environ/LITELLM_MASTER_KEY
tracing:
store:
type: clickhouse
url: os.environ/CLICKHOUSE_URL
retention_days: 14

View file

@ -0,0 +1,2 @@
httpx==0.28.1
pydantic==2.13.4

View file

@ -0,0 +1,172 @@
# This file was autogenerated by uv via the following command:
# uv pip compile --universal --python-version 3.13 --generate-hashes --no-emit-index-url deploy/lens/requirements.in -o deploy/lens/requirements.lock
annotated-types==0.8.0 \
--hash=sha256:13b2beaad985e05e2d6407ee4c4f35590b11f8d693a258a561055cac8f64cab7 \
--hash=sha256:f072f4d804ea359e4eaf198b1af7a8b0943881a87f31bb764f8bf219bb9419e0
# via pydantic
anyio==4.15.1 \
--hash=sha256:6152fdbbf9a77fdec97731721bebf7c4c44f7c29b424b0065826173efc7ed101 \
--hash=sha256:9f28306018cbd6d329e64a36d58256edff76dd996fe423bc957326e578b82a94
# via httpx
certifi==2026.7.22 \
--hash=sha256:62f22742b58a1a33014a2b6b706588a8d7e2a88ae7bd1a6ebe8c992928483775 \
--hash=sha256:741e2c3b351ddf169a738da9f2c048608ff7f2c5cc02f1ebc6b118bb090d5d55
# via
# httpcore
# httpx
h11==0.16.0 \
--hash=sha256:4e35b956cf45792e4caa5885e69fba00bdbc6ffafbfa020300e549b208ee5ff1 \
--hash=sha256:63cf8bbe7522de3bf65932fda1d9c2772064ffb3dae62d55932da54b31cb6c86
# via httpcore
httpcore==1.0.9 \
--hash=sha256:2d400746a40668fc9dec9810239072b40b4484b640a8c38fd654a024c7a1bf55 \
--hash=sha256:6e34463af53fd2ab5d807f399a9b45ea31c3dfa2276f15a2c3f00afff6e176e8
# via httpx
httpx==0.28.1 \
--hash=sha256:75e98c5f16b0f35b567856f597f06ff2270a374470a5c2392242528e3e3e42fc \
--hash=sha256:d909fcccc110f8c7faf814ca82a9a4d816bc5a6dbfea25d6591d6985b8ba59ad
# via -r deploy/lens/requirements.in
idna==3.20 \
--hash=sha256:a7db850025b95ded1eae8a46181a1a6c56c92c96f0e2b005d9ff8dc0210cab44 \
--hash=sha256:ab7ae7122974553370f0bdb919e1a960b2cd1bc1ef0276416d896db81c14582c
# via
# anyio
# httpx
pydantic==2.13.4 \
--hash=sha256:45a282cde31d808236fd7ea9d919b128653c8b38b393d1c4ab335c62924d9aba \
--hash=sha256:c40756b57adaa8b1efeeced5c196f3f3b7c435f90e84ea7f443901bec8099ef6
# via -r deploy/lens/requirements.in
pydantic-core==2.46.4 \
--hash=sha256:00c603d540afdd6b80eb39f078f33ebd46211f02f33e34a32d9f053bba711de0 \
--hash=sha256:0186750b482eefa11d7f435892b09c5c606193ef3375bcf94aa00ae6bfb66262 \
--hash=sha256:041bde0a48fd37cf71cab1c9d56d3e8625a3793fef1f7dd232b3ff37e978ecda \
--hash=sha256:0c563b08bca408dc7f65f700633d8442fffb2421fc47b8101377e9fd65051ff0 \
--hash=sha256:0cbe8b01f948de4286c74cdd6c667aceb38f5c1e26f0693b3983d9d74887c65e \
--hash=sha256:0ce40cd7b21210e99342afafbd4d0f76d784eb5b1d60f3bdc566be4983c6c73b \
--hash=sha256:0e96592440881c74a213e5ad528e2b24d3d4f940de2766bed9010ab1d9e51594 \
--hash=sha256:10e17cbb10a330363733efc4d7c4d0dd827ac0909b8f6a6542298fed1ea62f29 \
--hash=sha256:133878133d271ade3d41d1bfb2a45ec38dbdbda40bc065921c6b04e4630127e2 \
--hash=sha256:14d4edf427bdcf950a8a02d7cb44a08614388dd6e1bdcbf4f67504fa7887da9c \
--hash=sha256:14f4c5d6db102bd796a627bbb3a17b4cf4574b9ae861d8b7c9a9661c6dd3362d \
--hash=sha256:17299feefe090f2caa5b8e37222bb5f663e4935a8bfa6931d4102e5df1a9f398 \
--hash=sha256:184c081504d17f1c1066e430e117142b2c77d9448a97f7b65c6ac9fd9aee238d \
--hash=sha256:18e5ceec2ab67e6d5f1a9085e5a24c9c4e2ac4545730bfe668680bca05e555f3 \
--hash=sha256:19e51f073cd3df251856a8a4189fbdf1de4012c3ebacfb1884f94f1eb406079f \
--hash=sha256:1a7dd0b3ee80d90150e3495a3a13ac34dbcbfd4f012996a6a1d8900e91b5c0fb \
--hash=sha256:1d8ba486450b14f3b1d63bc521d410ec7565e52f887b9fb671791886436a42f7 \
--hash=sha256:2108ba5c1c1eca18030634489dc544844144ee36357f2f9f780b93e7ddbb44b5 \
--hash=sha256:228ee9bae8bef5b1e97ec58302f80357c37199e0d0a99174e138d28e6957b9d9 \
--hash=sha256:23ace664830ee0bfe014a0c7bc248b1f7f25ed7ad103852c317624a1083af462 \
--hash=sha256:2412e734dcb48da14d4e4006b82b46b74f2518b8a26ee7e58c6844a6cd6d03c4 \
--hash=sha256:29c61fc04a3d840155ff08e475a04809278972fe6aef51e2720554e96367e34b \
--hash=sha256:2f84c03c8607173d16b5a854ec68a2f9079ae03237a54fb506d13af47e1d018d \
--hash=sha256:3009f12e4e90b7f88b4f9adb1b0c4a3d58fe7820f3238c190047209d148026df \
--hash=sha256:3245406455a5d98187ec35530fd772b1d799b26667980872c8d4614991e2c4a2 \
--hash=sha256:3447661d99f75a3683a4cf5c87da72f2161964611864dbbeac7fbb118bb4bfc0 \
--hash=sha256:372429a130e469c9cd698925ce5fc50940b7a1336b0d82038e63d5bbc4edc519 \
--hash=sha256:395aebd9183f9d112f569aeb5b2214d1a10a33bec8456447f7fbdfa51d38d4cd \
--hash=sha256:3a233125ac121aa3ffba9a2b59edfc4a985a76092dc8279586ab4b71390875e7 \
--hash=sha256:3be77f45df024d789a672ae34f8b06fb346c4f9f46ea714956660ea4862e89ac \
--hash=sha256:3bf92c5d0e00fefaab325a4d27828fe6b6e2a21848686b5b60d2d9eeb09d76c6 \
--hash=sha256:3ecbc122d18468d06ca279dc26a8c2e2d5acb10943bb35e36ae92096dc3b5565 \
--hash=sha256:3fb702cd90b0446a3a1c5e470bfa0dd23c0233b676a9099ddcc964fa6ca13898 \
--hash=sha256:428e04521a40150c85216fc8b85e8d39fece235a9cf5e383761238c7fa9b96fb \
--hash=sha256:432c179df7874eeb73307aad2df0755e1ae0efa61ff0ea89b93e194411ae3928 \
--hash=sha256:4a05d69cba51d852c5c3e92758653245a50c0b646ced0cf05bd793ed592839d6 \
--hash=sha256:4c63ebc82684aa89d9a3bcbd13d515b3be44250dc68dd3bd81526c1cb31286c3 \
--hash=sha256:4fc73cb559bdb54b1134a706a2802a4cddd27a0633f5abb7e53056268751ac6a \
--hash=sha256:4fcbe087dbc2068af7eda3aa87634eba216dbda64d1ae73c8684b621d33f6596 \
--hash=sha256:56cb4851bcaf3d117eddcef4fe66afd750a50274b0da8e22be256d10e5611987 \
--hash=sha256:5855698a4856556d86e8e6cd8434bc3ac0314ee8e12089ae0e143f64c6256e4e \
--hash=sha256:5a4330cdbc57162e4b3aa303f588ba752257694c9c9be3e7ebb11b4aca659b5d \
--hash=sha256:5b712b53160b79a5850310b912a5ef8e57e56947c8ad690c227f5c9d7e561712 \
--hash=sha256:5d5902252db0d3cedf8d4a1bc68f70eeb430f7e4c7104c8c476753519b423008 \
--hash=sha256:617d7e2ca7dcb8c5cf6bcb8c59b8832c94b36196bbf1cbd1bfb56ed341905edd \
--hash=sha256:62f875393d7f270851f20523dd2e29f082bcc82292d66db2b64ea71f64b6e1c1 \
--hash=sha256:633147d34cf4550417f12e2b1a0383973bdf5cdfde212cb09e9a581cf10820be \
--hash=sha256:66ce7632c22d837c95301830e111ad0128a32b8207533b60896a96c4915192ea \
--hash=sha256:6b3ace8194b0e5204818c92802dcdca7fc6d88aabbb799d7c795540d9cd6d292 \
--hash=sha256:6f2eeda33a839975441c86a4119e1383c50b47faf0cbb5176985565c6bb02c33 \
--hash=sha256:7027560ee92211647d0d34e3f7cd6f50da56399d26a9c8ad0da286d3869a53f3 \
--hash=sha256:7283d57845ecf5a163403eb0702dfc220cc4fbdd18919cb5ccea4f95ee1cdab4 \
--hash=sha256:7a5f930472650a82629163023e630d160863fce524c616f4e5186e5de9d9a49b \
--hash=sha256:7bfb192b3f4b9e8a89b6277b6ce787564f62cfd272055f6e685726b111dc7826 \
--hash=sha256:811ff8e9c313ab425368bcbb36e5c4ebd7108c2bbf4e4089cfbb0b01eff63fac \
--hash=sha256:8233f2947cf85404441fd7e0085f53b10c93e0ee78611099b5c7237e36aacbf7 \
--hash=sha256:82cf5301172168103724d49a1444d3378cb20cdee30b116a1bd6031236298a5d \
--hash=sha256:8358a950c8909158e3df31538a7e4edc2d7265a7c54b47f0864d9e5bae9dcebf \
--hash=sha256:85bb3611ff1802f3ee7fdd7dbff26b56f343fb432d57a4728fdd49b6ef35e2f4 \
--hash=sha256:86e1a4418c6cd97d60c95c71164158eaf7324fae7b0923264016baa993eba6fc \
--hash=sha256:8b9bab013d1c7a79d3501ff86d0bc9c31bf587db4551677b96bec07df78c6b15 \
--hash=sha256:8c5dac79fa1614d1e06ca695109c6105923bd9c7d1d6c918d4e637b7e6b32fd3 \
--hash=sha256:8d0820e8192167f80d88d64038e609c31452eeca865b4e1d9950a27a4609b00b \
--hash=sha256:8daafc69c93ee8a0204506a3b6b30f586ef54028f52aeeeb5c4cfc5184fd5914 \
--hash=sha256:9037063db01f09b09e237c282b6792bd4da634b5402c4e7f0c61effed7701a04 \
--hash=sha256:905a0ed8ea6f2d61c1738835f99b699348d7857379083e5fc497fa0c967a407c \
--hash=sha256:90884113d8b48f760e9587002789ddd741e76ab9f89518cd1e43b1f1a52ec44b \
--hash=sha256:91a06d2e259ecfbd8c901d70c3c507900458498142b3026a296b7de4d1322cc9 \
--hash=sha256:926c9541b14b12b1681dca8a0b75feb510b06c6341b70a8e500c2fdcff837cce \
--hash=sha256:9401557acd873c3a7f3eb9383edef8ac4968f9510e340f4808d427e75667e7b4 \
--hash=sha256:9551187363ffc0de2a00b2e47c25aeaeb1020b69b668762966df15fc5659dd5a \
--hash=sha256:962ccbab7b642487b1d8b7df90ef677e03134cf1fd8880bf698649b22a69371f \
--hash=sha256:97e7cf2be5c77b7d1a9713a05605d49460d02c6078d38d8bef3cbe323c548424 \
--hash=sha256:9aa768456404a8bf48a4406685ac2bec8e72b62c69313734fa3b73cf33b3a894 \
--hash=sha256:9bc519fbf2b7578398853d815009ae5e4d4603d12f4e3f91da8c06852d3da3e9 \
--hash=sha256:9d56801be94b86a9da183e5f3766e6310752b99ff647e38b09a9500d88e46e76 \
--hash=sha256:9f444c499b3eefd3a92e348059471ea0c3a6e303d9c1cec09fa748fd9f895201 \
--hash=sha256:9fa8ae11da9e2b3126c6426f147e0fba88d96d65921799bb30c6abd1cb2c97fb \
--hash=sha256:a0f62d0a58f4e7da165457e995725421e0064f2255d8eccebc49f41bbc23b109 \
--hash=sha256:a396dcc17e5a0b164dbe026896245a4fa9ff402edca1dff0be3d53a517f74de4 \
--hash=sha256:aaa2a54443eff1950ba5ddc6b6ccda0d9c84a364276a62f969bdf2a390650848 \
--hash=sha256:ad785e92e6dc634c21555edc8bd6b64957ab844541bcb96a1366c202951ae526 \
--hash=sha256:af8244b2bef6aaad6d92cda81372de7f8c8d36c9f0c3ea36e827c60e7d9467a0 \
--hash=sha256:b078afbc25f3a1436c7a1d2cd3e322497ee99615ba97c563566fdf46aff1ee01 \
--hash=sha256:b2f69dec1725e79a012d920df1707de5caf7ed5e08f3be4435e25803efc47458 \
--hash=sha256:b8458003118a712e66286df6a707db01c52c0f52f7db8e4a38f0da1d3b94fc4e \
--hash=sha256:bb63e0198ca18aad131c089b9204c23079c3afa95487e561f4c522d519e55aba \
--hash=sha256:bfec22eab3c8cc2ceec0248aec886624116dc079afa027ecc8ad4a7e62010f8a \
--hash=sha256:c1747f85cee84c26985853c6f3d9bd3e75da5212912443fa111c113b9c246f39 \
--hash=sha256:c1b3f518abeca3aa13c712fd202306e145abf59a18b094a6bafb2d2bbf59192c \
--hash=sha256:c50f2528cf200c5eed56faf3f4e22fcd5f38c157a8b78576e6ba3168ec35f000 \
--hash=sha256:c68fcd102d71ea85c5b2dfac3f4f8476eff42a9e078fd5faefff6d145063536b \
--hash=sha256:c7a7bd4e39e8e4c12c39cd480356842b6a8a06e41b23a55a5e3e191718838ddf \
--hash=sha256:c94f0688e7b8d0a67abf40e57a7eaaecd17cc9586706a31b76c031f63df052b4 \
--hash=sha256:cbaf13819775b7f769bf4a1f066cb6df7a28d4480081a589828ef190226881cd \
--hash=sha256:cd2213145bcc2ba85884d0ac63d222fece9209678f77b9b4d76f054c561adb28 \
--hash=sha256:ce5c1d2a8b27468f433ca974829c44060b8097eedc39933e3c206a90ee49c4a9 \
--hash=sha256:d396ec2b979760aaf3218e76c24e65bd0aca24983298653b3a9d7a45f9e47b30 \
--hash=sha256:d51026d73fcfd93610abc7b27789c26b313920fcfb20e27462d74a7f8b06e983 \
--hash=sha256:d80ee3d731373b24cebbc10d689ca4ee1875caf0d5703a245db18efd4dd37fc1 \
--hash=sha256:d995260fdf4e1db774581b4900e0f832abe3c7c84996726bbc161b19c8f29e76 \
--hash=sha256:da4b951fe36dc7c3a1ccb4e3cd1747c3542b8c9ceede8fc86cae054e764485f5 \
--hash=sha256:daa27d92c36f24388fe3ad306b174781c747627f134452e4f128ea00ce1fe8c4 \
--hash=sha256:db06ffe51636ffe9ca531fe9023dd64bdd794be8754cb5df57c5498ae5b518a7 \
--hash=sha256:e0d65b8c354be7fb5f720c3caa8bc940bc2d20ce749c8e06135f07f8ed95dd7c \
--hash=sha256:e68b7a074f65a2fd746c52a7ce6142ab7006074ac269ace0c25cd8ba171f8066 \
--hash=sha256:e739fee756ba1010f8bcccb534252e85a35fe45ae92c295a06059ce58b74ccd3 \
--hash=sha256:e846ae7835bf0703ae43f534ab79a867146dadd59dc9ca5c8b53d5c8f7c9ef02 \
--hash=sha256:e9c26f834c65f5752f3f06cb08cb86a913ceb7274d0db6e267808a708b46bc89 \
--hash=sha256:ea793e075b70290d89d8142074262885d3f7da19634845135751bd6344f73b50 \
--hash=sha256:f027324c56cd5406ca49c124b0db10e56c69064fec039acc571c29020cc87c76 \
--hash=sha256:f13a646d65d09fbf1bc6b3a9635d30095c8e7e5cc419ff35ecc563c5fd04cd49 \
--hash=sha256:f47286a97f0bc9b8859519809077b91b2cefe4ae47fcbf5e466a009c1c5d742b \
--hash=sha256:f747929cf940cddb5b3668a390056ddd5ba2e5010615ea2dcf4f9c4f3ab8791d \
--hash=sha256:f99626688942fb746e545232e7726926f3be91b5975f8b55327665fafda991c7 \
--hash=sha256:f9fa868638bf362d3d138ea55829cefb3d5f4b0d7f142234382a15e2485dbec4 \
--hash=sha256:fbdb89b3e1c94a30cc5edfce477c6e6a5dc4d8f84665b455c27582f211a1c72c \
--hash=sha256:fc010ab034c8c7452522748bf937df58020d256ccae0874463d1f4d01758af8e \
--hash=sha256:fc3e9034a63de20e15e8ade85358bc6efc614008cab72898b4b4952bea0509ff \
--hash=sha256:fd8b3d9fd264be37976686c7f65cd52a83f5e84f4bfd2adf9c1d469676bbb6ae
# via pydantic
typing-extensions==4.16.0 \
--hash=sha256:481caa481374e813c1b176ada14e97f1f67a4539ce9cfeb3f350d78d6370c2e8 \
--hash=sha256:dc983d19a509c94dba722ee6abd33940f7c05a89e243c47e907eb4db6f1a43e5
# via
# anyio
# pydantic
# pydantic-core
# typing-inspection
typing-inspection==0.4.4 \
--hash=sha256:547274fa6b0a561ccf549cc9524b999a578e737d015d8709d021f9d0d13bea47 \
--hash=sha256:65b8397ba37ccbce054456aaccddfc91e6e3083c92824df348d96ca832f3f147
# via pydantic

91
deploy/lens/stack.yaml Normal file
View file

@ -0,0 +1,91 @@
name: litellm-lens
services:
litellm:
image: ghcr.io/berriai/litellm:${LITELLM_VERSION:?Set LITELLM_VERSION to a published release, without the v prefix}
entrypoint:
- python3
- -c
- |
import os, sys
from urllib.parse import quote
postgres_password = quote(os.environ["POSTGRES_PASSWORD"], safe="")
clickhouse_password = quote(os.environ["CLICKHOUSE_PASSWORD"], safe="")
os.environ["DATABASE_URL"] = f"postgresql://litellm:{postgres_password}@db:5432/litellm"
os.environ["CLICKHOUSE_URL"] = f"http://default:{clickhouse_password}@clickhouse:8123"
os.execv("docker/prod_entrypoint.sh", ["docker/prod_entrypoint.sh", *sys.argv[1:]])
command: ["--config", "/app/lens-config.yaml", "--port", "4000"]
environment:
LITELLM_MASTER_KEY: ${LITELLM_MASTER_KEY:?Set a strong master key}
LITELLM_SALT_KEY: ${LITELLM_SALT_KEY:?Set a permanent encryption key and keep it across upgrades}
POSTGRES_PASSWORD: ${POSTGRES_PASSWORD:?Set a permanent database password}
STORE_MODEL_IN_DB: "True"
CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD:?Set a permanent ClickHouse password}
LENS_WORKER_IMAGE: ghcr.io/berriai/litellm-lens-worker:v${LITELLM_VERSION}
volumes:
- ./config.yaml:/app/lens-config.yaml:ro
ports:
- "127.0.0.1:${LITELLM_PORT:-4000}:4000"
networks: [proxy, storage]
depends_on:
db:
condition: service_healthy
clickhouse:
condition: service_healthy
restart: unless-stopped
lens-worker:
profiles: [lens]
image: ghcr.io/berriai/litellm-lens-worker:v${LITELLM_VERSION}
environment:
LITELLM_URL: http://litellm:4000
LENS_WORKER_TOKEN: ${LENS_WORKER_TOKEN:-}
depends_on: [litellm]
networks: [proxy]
restart: unless-stopped
read_only: true
tmpfs:
- /tmp:rw,noexec,nosuid,size=${LENS_WORKER_TMP_SIZE:-1g}
cap_drop: [ALL]
security_opt: [no-new-privileges:true]
db:
image: postgres:16
environment:
POSTGRES_DB: litellm
POSTGRES_USER: litellm
POSTGRES_PASSWORD: ${POSTGRES_PASSWORD}
networks: [storage]
volumes:
- postgres_data:/var/lib/postgresql/data
healthcheck:
test: ["CMD-SHELL", "pg_isready -U litellm -d litellm"]
interval: 5s
timeout: 5s
retries: 20
restart: unless-stopped
clickhouse:
image: clickhouse/clickhouse-server:26.9.6.6
environment:
CLICKHOUSE_USER: default
CLICKHOUSE_PASSWORD: ${CLICKHOUSE_PASSWORD}
CLICKHOUSE_DEFAULT_ACCESS_MANAGEMENT: "1"
volumes:
- clickhouse_data:/var/lib/clickhouse
healthcheck:
test: ["CMD", "clickhouse-client", "--user", "default", "--password", "${CLICKHOUSE_PASSWORD}", "--query", "SELECT 1"]
interval: 5s
timeout: 5s
retries: 20
restart: unless-stopped
networks: [storage]
networks:
proxy:
storage:
internal: true
volumes:
postgres_data:
clickhouse_data:

View file

@ -0,0 +1,44 @@
services:
litellm:
image: ${LITELLM_IMAGE:?Set the native-enabled gateway image}
environment:
LITELLM_ADMIN_AGENT_URL: http://liteadmin:10000
ADMIN_AGENT_SERVICE_TOKEN: ${ADMIN_AGENT_SERVICE_TOKEN:?Set a shared worker token}
PROXY_BASE_URL: ${LITELLM_PUBLIC_URL:?Set the existing HTTPS gateway URL}
liteadmin:
image: ${LITELLM_IMAGE:?Set the same native-enabled image used by the gateway}
command: ["--admin-agent"]
restart: unless-stopped
init: true
read_only: true
cap_drop: [ALL]
security_opt: [no-new-privileges:true]
stop_grace_period: 75s
environment:
CONNECTION_AUTH_MODE: native
LITELLM_BASE_URL: ${LITELLM_PUBLIC_URL:?Set the existing HTTPS gateway URL}
LITELLM_MODEL: ${LITELLM_ADMIN_MODEL:?Set a gateway model with tool support}
SLACK_BOT_TOKEN: ${SLACK_BOT_TOKEN:?Install the Slack app}
SLACK_APP_TOKEN: ${SLACK_APP_TOKEN:?Enable Socket Mode}
SLACK_WORKSPACE_ID: ${SLACK_WORKSPACE_ID:?Set the Slack workspace ID}
ADMIN_AGENT_SERVICE_TOKEN: ${ADMIN_AGENT_SERVICE_TOKEN:?Set a shared worker token}
CREDENTIAL_ENCRYPTION_KEY: ${CREDENTIAL_ENCRYPTION_KEY:?Set a persistent Fernet key}
STATE_DB: /var/data/events.sqlite3
ADMIN_READ_ONLY: ${ADMIN_READ_ONLY:-false}
OPENAI_AGENTS_DISABLE_TRACING: "1"
volumes:
- liteadmin_state:/var/data
tmpfs:
- /tmp:rw,noexec,nosuid,size=64m
healthcheck:
test: ["CMD", "/opt/liteadmin/bin/python", "-c", "import urllib.request; urllib.request.urlopen('http://127.0.0.1:10000/readyz', timeout=3)"]
interval: 30s
timeout: 5s
start_period: 30s
depends_on:
litellm:
condition: service_healthy
volumes:
liteadmin_state:

View file

@ -113,6 +113,8 @@ RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \
sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
FROM $LITELLM_RUNTIME_IMAGE AS runtime
ARG LITELLM_RELEASE_TAG=""
ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG}
USER root

View file

@ -122,6 +122,8 @@ RUN sed -i 's/\r$//' docker/entrypoint.sh && chmod +x docker/entrypoint.sh && \
sed -i 's/\r$//' docker/prod_entrypoint.sh && chmod +x docker/prod_entrypoint.sh
FROM $LITELLM_RUNTIME_IMAGE AS runtime
ARG LITELLM_RELEASE_TAG=""
ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG}
WORKDIR /app
USER root

View file

@ -1,5 +1,11 @@
#!/bin/sh
if [ "$1" = "--admin-agent" ]; then
shift
export CONNECTION_AUTH_MODE=native
exec /opt/liteadmin/bin/litellm-admin-agent --web "$@"
fi
case "$USE_DDTRACE" in
[Tt][Rr][Uu][Ee])
export DD_TRACE_OPENAI_ENABLED="False"

View file

@ -6,6 +6,7 @@ from litellm_enterprise.enterprise_callbacks.send_emails.endpoints import (
from . import ui_crud_endpoints # side-effect: registers extra UI settings
from .audit_logging_endpoints import router as audit_logging_router
from .liteadmin import router as liteadmin_router
from .management_endpoints import management_endpoints_router
from .utils import _should_block_robots
@ -14,6 +15,7 @@ __all__ = ["router", "ui_crud_endpoints"]
router = APIRouter()
router.include_router(email_events_router)
router.include_router(audit_logging_router)
router.include_router(liteadmin_router)
router.include_router(management_endpoints_router)

View file

@ -0,0 +1,283 @@
from __future__ import annotations
import hashlib
import hmac
import html
import os
import re
import secrets
from collections.abc import Awaitable, Callable
from dataclasses import dataclass
from datetime import datetime, timedelta, timezone
from typing import Annotated, Final
from urllib.parse import urlencode, urlsplit
import httpx
from fastapi import APIRouter, Depends, HTTPException, Request
from fastapi.responses import HTMLResponse, RedirectResponse, Response
from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.proxy._experimental.mcp_server.oauth_utils import get_request_base_url
from litellm.proxy._types import LiteLLM_UserTable, LitellmUserRoles, UserAPIKeyAuth
from litellm.types.proxy.auth.auth_checks import UserNotFoundError
router: Final = APIRouter()
_PREFIX: Final = "/liteadmin/slack/connect/"
_COOKIE: Final = "__Host-litellm-slack-connect-"
_HEADERS: Final = {
"Cache-Control": "no-store",
"Referrer-Policy": "same-origin",
"X-Frame-Options": "DENY",
"X-Content-Type-Options": "nosniff",
"Content-Security-Policy": "default-src 'none'; style-src 'unsafe-inline'; form-action 'self'; frame-ancestors 'none'; base-uri 'none'",
}
class LinkDetails(BaseModel):
model_config = ConfigDict(frozen=True, strict=True, extra="forbid")
workspace_id: str = Field(min_length=1, max_length=64)
slack_user_id: str = Field(min_length=1, max_length=64)
email: str = Field(min_length=1, max_length=320)
class AdminSession(BaseModel):
model_config = ConfigDict(frozen=True)
user_id: str
credential: SecretStr
expires_at: float
@dataclass(frozen=True, slots=True)
class NativeAdminContext:
worker_url: str
service_token: SecretStr
client: httpx.AsyncClient
session_user: Callable[[Request], Awaitable[str | None]]
load_user: Callable[[str], Awaitable[LiteLLM_UserTable | None]]
mint_session: Callable[[LiteLLM_UserTable], AdminSession]
async def worker_request(self, token: str, session: AdminSession | None = None) -> httpx.Response:
if re.fullmatch(r"[A-Za-z0-9_-]{43}", token) is None:
raise HTTPException(410, "Connection link expired. Send connect in Slack for a new link")
try:
response: Final = await self.client.request(
"GET" if session is None else "POST",
f"{self.worker_url}/internal/liteadmin/links/{token}",
headers={"X-LiteLLM-Admin-Agent-Token": self.service_token.get_secret_value()},
json=None
if session is None
else {
"user_id": session.user_id,
"credential": session.credential.get_secret_value(),
"expires_at": session.expires_at,
},
timeout=15,
follow_redirects=False,
)
except httpx.HTTPError:
raise HTTPException(503, "LiteAdmin is temporarily unavailable") from None
if response.status_code == 410:
raise HTTPException(410, "Connection link expired. Send connect in Slack for a new link")
if response.status_code == 403:
raise HTTPException(403, "Connect your own active LiteLLM proxy-admin account with the same email as Slack")
if response.status_code != 200:
raise HTTPException(503, "LiteAdmin could not verify this connection")
return response
async def details(self, token: str) -> LinkDetails:
response: Final = await self.worker_request(token)
try:
return LinkDetails.model_validate_json(response.content)
except ValidationError:
raise HTTPException(503, "LiteAdmin could not verify this connection") from None
async def admin(self, user_id: str, details: LinkDetails) -> LiteLLM_UserTable:
user: Final = await self.load_user(user_id)
if (
user is None
or user.user_role != LitellmUserRoles.PROXY_ADMIN.value
or not user.user_email
or user.user_email.strip().casefold() != details.email.strip().casefold()
):
raise HTTPException(403, "Connect your own active LiteLLM proxy-admin account with the same email as Slack")
return user
def _page(title: str, body: str) -> HTMLResponse:
return HTMLResponse(
f'<!doctype html><html lang="en"><meta charset="utf-8">'
f'<meta name="viewport" content="width=device-width,initial-scale=1"><title>{html.escape(title)}</title>'
"<style>body{font:17px system-ui;color:#18252f;max-width:560px;margin:10vh auto;padding:24px}"
"p{line-height:1.6}button{font:inherit;border:0;border-radius:8px;padding:14px 20px;background:#5b3fd1;"
"color:white;cursor:pointer}small{color:#556}</style>"
f"<main><h1>{html.escape(title)}</h1>{body}</main></html>",
headers=_HEADERS,
)
def _cookie_name(token: str) -> str:
return _COOKIE + hashlib.sha256(token.encode()).hexdigest()[:16]
async def _session_user(request: Request) -> str | None:
from litellm.proxy._experimental.mcp_server.byok_oauth_endpoints import (
get_authenticated_browser_user_id,
)
return await get_authenticated_browser_user_id(request)
async def _load_user(user_id: str) -> LiteLLM_UserTable | None:
from litellm.proxy.auth.auth_checks import get_user_object
from litellm.proxy.proxy_server import prisma_client, user_api_key_cache
if prisma_client is None:
raise HTTPException(503, "LiteAdmin requires a database")
try:
return await get_user_object(
user_id=user_id,
prisma_client=prisma_client,
user_api_key_cache=user_api_key_cache,
user_id_upsert=False,
check_db_only=True,
)
except UserNotFoundError:
return None
except Exception:
raise HTTPException(503, "LiteAdmin could not verify your current permissions") from None
def mint_admin_session(user: LiteLLM_UserTable) -> AdminSession:
from litellm.proxy.auth.auth_checks import LITELLM_SESSION_TOKEN_PREFIX
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_bearer_token
expires: Final = datetime.now(timezone.utc) + timedelta(hours=24)
auth: Final = UserAPIKeyAuth(
token="liteadmin-" + secrets.token_urlsafe(24),
key_name="LiteAdmin Slack",
key_alias="LiteAdmin Slack",
user_id=user.user_id,
user_role=LitellmUserRoles.PROXY_ADMIN,
models=TypeAdapter(list[str]).validate_python(user.model_dump().get("models", [])),
expires=expires,
is_session_token=True,
)
return AdminSession(
user_id=user.user_id,
credential=SecretStr(
encrypt_bearer_token(auth.model_dump_json(exclude_none=True), LITELLM_SESSION_TOKEN_PREFIX)
),
expires_at=expires.timestamp(),
)
def validate_native_configuration(
worker_url: str, service_token: str, enterprise: bool, database_available: bool
) -> None:
if not worker_url:
raise HTTPException(404, "LiteAdmin Slack is not enabled")
if not enterprise:
raise HTTPException(403, "LiteAdmin Slack requires LiteLLM Enterprise")
if not database_available:
raise HTTPException(503, "LiteAdmin requires a database")
try:
parsed: Final = urlsplit(worker_url)
port: Final = parsed.port
except ValueError:
raise HTTPException(503, "LiteAdmin worker configuration is invalid") from None
if (
parsed.scheme not in {"http", "https"}
or not parsed.hostname
or port == 0
or parsed.username
or parsed.password
or parsed.path
or parsed.query
or parsed.fragment
or len(service_token) < 32
or any(character.isspace() for character in service_token)
):
raise HTTPException(503, "LiteAdmin worker configuration is invalid")
async def native_admin_context() -> NativeAdminContext:
from litellm.proxy.proxy_server import premium_user, prisma_client
worker_url: Final = os.getenv("LITELLM_ADMIN_AGENT_URL", "").rstrip("/")
service_token: Final = os.getenv("ADMIN_AGENT_SERVICE_TOKEN", "")
validate_native_configuration(worker_url, service_token, premium_user is True, prisma_client is not None)
client: Final = get_async_httpx_client(
llm_provider="liteadmin_native", params={"timeout": 15.0, "follow_redirects": False}
).client
return NativeAdminContext(
worker_url, SecretStr(service_token), client, _session_user, _load_user, mint_admin_session
)
@router.get(_PREFIX + "{token}", include_in_schema=False, response_class=HTMLResponse)
async def connect_page(
request: Request,
token: str,
context: Annotated[NativeAdminContext, Depends(native_admin_context)],
) -> Response:
details: Final = await context.details(token)
base_url: Final = get_request_base_url(request)
parsed_base: Final = urlsplit(base_url)
if parsed_base.scheme != "https":
raise HTTPException(400, "LiteAdmin account connections require HTTPS")
user_id: Final = await context.session_user(request)
if user_id is None:
return RedirectResponse(
base_url + "/sso/key/generate?" + urlencode({"return_to": parsed_base.path + _PREFIX + token}),
status_code=303,
headers=_HEADERS,
)
await context.admin(user_id, details)
csrf: Final = secrets.token_urlsafe(32)
page: Final = _page(
"Connect LiteAdmin to Slack",
f"<p>Connect <strong>{html.escape(details.email)}</strong> to LiteAdmin in your Slack workspace?</p>"
"<p>Model requests and administrative actions will use your own LiteLLM account and current permissions</p>"
f'<form method="post"><input type="hidden" name="csrf" value="{csrf}">'
'<button type="submit">Connect account</button></form>'
"<p><small>This connection lasts 24 hours. Send disconnect in Slack to remove the saved session</small></p>",
)
page.set_cookie(_cookie_name(token), csrf, max_age=600, secure=True, httponly=True, samesite="strict", path="/")
return page
@router.post(_PREFIX + "{token}", include_in_schema=False, response_class=HTMLResponse)
async def connect_account(
request: Request,
token: str,
context: Annotated[NativeAdminContext, Depends(native_admin_context)],
) -> Response:
base_url: Final = get_request_base_url(request)
parsed_base: Final = urlsplit(base_url)
origin: Final = f"{parsed_base.scheme}://{parsed_base.netloc}"
if parsed_base.scheme != "https" or request.headers.get("Origin") != origin:
raise HTTPException(403, "Reopen your private Slack connection link")
if request.headers.get("Content-Type", "").split(";", 1)[0] != "application/x-www-form-urlencoded":
raise HTTPException(400, "Expected a connection form")
form: Final = await request.form(max_fields=1, max_files=0, max_part_size=1024)
supplied: Final = form.get("csrf")
expected: Final = request.cookies.get(_cookie_name(token), "")
if (
not isinstance(supplied, str)
or len(expected) != 43
or len(supplied) != 43
or not hmac.compare_digest(supplied.encode(), expected.encode())
):
raise HTTPException(403, "Reopen your private Slack connection link")
user_id: Final = await context.session_user(request)
if user_id is None:
raise HTTPException(401, "Your login expired. Reopen your private Slack connection link")
details: Final = await context.details(token)
user: Final = await context.admin(user_id, details)
await context.worker_request(token, context.mint_session(user))
page: Final = _page(
"Account connected", "<p>Return to Slack and ask LiteAdmin to list your teams or check a budget</p>"
)
page.delete_cookie(_cookie_name(token), path="/", secure=True, httponly=True, samesite="strict")
return page

View file

@ -57,6 +57,19 @@ spec:
imagePullPolicy: {{ .Values.image.pullPolicy }}
env:
{{- include "litellm.proxyEnv" . | nindent 12 }}
{{- if .Values.liteadmin.enabled }}
- name: LITELLM_ADMIN_AGENT_URL
value: {{ printf "http://%s-liteadmin:10000" (include "litellm.fullname" . | trunc 53 | trimSuffix "-") | quote }}
- name: ADMIN_AGENT_SERVICE_TOKEN
valueFrom:
secretKeyRef:
name: {{ required "liteadmin.existingSecret is required" .Values.liteadmin.existingSecret }}
key: ADMIN_AGENT_SERVICE_TOKEN
{{- if not (hasKey (default dict .Values.envVars) "PROXY_BASE_URL") }}
- name: PROXY_BASE_URL
value: {{ required "liteadmin.gatewayUrl is required" .Values.liteadmin.gatewayUrl | quote }}
{{- end }}
{{- end }}
{{- include "litellm.proxyMetricsEnv" . | nindent 12 }}
{{- if .Values.collector.enabled }}
{{- include "litellm.collectorEnv" . | nindent 12 }}

View file

@ -0,0 +1,112 @@
{{- if .Values.liteadmin.enabled }}
{{- $name := printf "%s-liteadmin" (include "litellm.fullname" . | trunc 53 | trimSuffix "-") }}
{{- $secret := required "liteadmin.existingSecret is required" .Values.liteadmin.existingSecret }}
apiVersion: apps/v1
kind: Deployment
metadata:
name: {{ $name }}
spec:
replicas: 1
strategy:
type: Recreate
selector:
matchLabels:
app.kubernetes.io/name: {{ $name }}
app.kubernetes.io/instance: {{ .Release.Name }}
template:
metadata:
labels:
app.kubernetes.io/name: {{ $name }}
app.kubernetes.io/instance: {{ .Release.Name }}
spec:
automountServiceAccountToken: false
terminationGracePeriodSeconds: 75
{{- with .Values.imagePullSecrets }}
imagePullSecrets:
{{- toYaml . | nindent 8 }}
{{- end }}
securityContext:
runAsUser: 10001
runAsGroup: 10001
fsGroup: 10001
runAsNonRoot: true
containers:
- name: liteadmin
image: "{{ .Values.image.repository }}:{{ .Values.image.tag | default .Chart.AppVersion }}"
imagePullPolicy: {{ .Values.image.pullPolicy }}
args: ["--admin-agent"]
securityContext:
allowPrivilegeEscalation: false
readOnlyRootFilesystem: true
capabilities:
drop: [ALL]
envFrom:
- secretRef:
name: {{ $secret }}
env:
- name: CONNECTION_AUTH_MODE
value: native
- name: LITELLM_BASE_URL
value: {{ required "liteadmin.gatewayUrl is required" .Values.liteadmin.gatewayUrl | quote }}
- name: LITELLM_MODEL
value: {{ required "liteadmin.model is required" .Values.liteadmin.model | quote }}
- name: STATE_DB
value: /var/data/events.sqlite3
- name: ADMIN_READ_ONLY
value: {{ .Values.liteadmin.readOnly | quote }}
- name: OPENAI_AGENTS_DISABLE_TRACING
value: "1"
ports:
- name: health
containerPort: 10000
readinessProbe:
httpGet:
path: /readyz
port: health
periodSeconds: 15
livenessProbe:
httpGet:
path: /healthz
port: health
periodSeconds: 30
resources:
{{- toYaml .Values.liteadmin.resources | nindent 12 }}
volumeMounts:
- name: state
mountPath: /var/data
- name: tmp
mountPath: /tmp
volumes:
- name: state
persistentVolumeClaim:
claimName: {{ $name }}
- name: tmp
emptyDir:
sizeLimit: 64Mi
---
apiVersion: v1
kind: Service
metadata:
name: {{ $name }}
spec:
type: ClusterIP
selector:
app.kubernetes.io/name: {{ $name }}
app.kubernetes.io/instance: {{ .Release.Name }}
ports:
- port: 10000
targetPort: health
---
apiVersion: v1
kind: PersistentVolumeClaim
metadata:
name: {{ $name }}
spec:
accessModes: [ReadWriteOnce]
{{- with .Values.liteadmin.storageClassName }}
storageClassName: {{ . | quote }}
{{- end }}
resources:
requests:
storage: {{ .Values.liteadmin.storageSize }}
{{- end }}

View file

@ -3,6 +3,20 @@
# Declare variables to be passed into your templates.
replicaCount: 1
liteadmin:
enabled: false
existingSecret: ""
gatewayUrl: ""
model: ""
readOnly: false
storageSize: 1Gi
storageClassName: ""
resources:
requests:
cpu: 100m
memory: 256Mi
limits:
memory: 1Gi
# numWorkers: 2
image:

View file

@ -471,6 +471,24 @@ Directory of the collector's unix socket, shared by the gateway and
collector containers through an emptyDir. Empty when the sidecar is off
or gateway.collector.address is a tcp://127.0.0.1:<port> address.
*/}}
{{- define "litellm.lensWorker.image" -}}
{{- if .Values.lensWorker.image.digest -}}
{{- if not (regexMatch "^sha256:[0-9a-f]{64}$" .Values.lensWorker.image.digest) -}}
{{- fail "lensWorker.image.digest must be sha256 followed by 64 lowercase hex characters" -}}
{{- end -}}
{{- printf "%s@%s" .Values.lensWorker.image.repository .Values.lensWorker.image.digest -}}
{{- else -}}
{{- $backendTag := .Values.backend.image.tag | default .Chart.AppVersion -}}
{{- $releaseTag := ternary (printf "v%s" $backendTag) $backendTag (regexMatch "^[0-9]" $backendTag) -}}
{{- $tag := .Values.lensWorker.image.tag | default $releaseTag -}}
{{- $repository := .Values.lensWorker.image.repository -}}
{{- if and (hasPrefix "sha-" $tag) (eq $repository "ghcr.io/berriai/litellm-lens-worker") -}}
{{- $repository = "ghcr.io/berriai/litellm-lens-worker-dev" -}}
{{- end -}}
{{- printf "%s:%s" $repository $tag -}}
{{- end -}}
{{- end -}}
{{- define "litellm.gateway.collectorSocketDir" -}}
{{- if and .Values.gateway.collector.enabled (hasPrefix "unix://" .Values.gateway.collector.address) -}}
{{- dir (trimPrefix "unix://" .Values.gateway.collector.address) -}}

View file

@ -57,6 +57,8 @@ spec:
containerPort: 4001
protocol: TCP
env:
- name: LENS_WORKER_IMAGE
value: {{ include "litellm.lensWorker.image" . | quote }}
{{- include "litellm.serverEnv" (dict "root" $ "component" .Values.backend) | nindent 12 }}
{{- if .Values.gateway.config.create }}
- name: CONFIG_FILE_PATH

View file

@ -0,0 +1,72 @@
{{- if .Values.lensWorker.enabled }}
apiVersion: apps/v1
kind: Deployment
metadata:
name: {{ include "litellm.fullname" . }}-lens-worker
labels:
{{- include "litellm.commonLabels" . | nindent 4 }}
app.kubernetes.io/component: lens-worker
spec:
replicas: {{ .Values.lensWorker.replicaCount }}
selector:
matchLabels:
app.kubernetes.io/instance: {{ .Release.Name }}
app.kubernetes.io/component: lens-worker
template:
metadata:
labels:
{{- include "litellm.commonLabels" . | nindent 8 }}
app.kubernetes.io/component: lens-worker
spec:
automountServiceAccountToken: false
{{- with .Values.imagePullSecrets }}
imagePullSecrets:
{{- toYaml . | nindent 8 }}
{{- end }}
securityContext:
runAsNonRoot: true
runAsUser: 65532
runAsGroup: 65532
fsGroup: 65532
seccompProfile:
type: RuntimeDefault
containers:
- name: lens-worker
image: {{ include "litellm.lensWorker.image" . | quote }}
imagePullPolicy: {{ .Values.lensWorker.image.pullPolicy }}
securityContext:
allowPrivilegeEscalation: false
readOnlyRootFilesystem: true
capabilities:
drop: [ALL]
env:
- name: LITELLM_URL
value: {{ .Values.lensWorker.url | default (printf "http://%s:%v" (include "litellm.backend.fullname" .) .Values.backend.service.port) | quote }}
- name: LENS_WORKER_TOKEN
valueFrom:
secretKeyRef:
name: {{ required "lensWorker.tokenSecret.name must reference a Lens worker token" .Values.lensWorker.tokenSecret.name | quote }}
key: {{ .Values.lensWorker.tokenSecret.key | quote }}
resources:
{{- toYaml .Values.lensWorker.resources | nindent 12 }}
volumeMounts:
- name: tmp
mountPath: /tmp
volumes:
- name: tmp
emptyDir:
medium: Memory
sizeLimit: {{ .Values.lensWorker.tmpSizeLimit }}
{{- with .Values.lensWorker.nodeSelector }}
nodeSelector:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.lensWorker.tolerations }}
tolerations:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- with .Values.lensWorker.affinity }}
affinity:
{{- toYaml . | nindent 8 }}
{{- end }}
{{- end }}

View file

@ -0,0 +1,174 @@
suite: Lens worker release and credentials
templates:
- lens/deployment.yaml
- backend/deployment.yaml
- gateway/configmap.yaml
values:
- ./values/required.yaml
tests:
- it: installs the development package for a source commit
template: lens/deployment.yaml
set:
backend.image.tag: sha-0123456789abcdef
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker-dev:sha-0123456789abcdef
- it: advertises the development package for standalone source workers
template: backend/deployment.yaml
set:
backend.image.tag: sha-0123456789abcdef
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LENS_WORKER_IMAGE
value: ghcr.io/berriai/litellm-lens-worker-dev:sha-0123456789abcdef
- it: preserves an explicit private source image repository
template: lens/deployment.yaml
set:
backend.image.tag: sha-0123456789abcdef
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
lensWorker.image.repository: registry.example/lens-worker
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: registry.example/lens-worker:sha-0123456789abcdef
- it: pins the worker to its approved digest even when its tag changes
template: lens/deployment.yaml
set:
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
lensWorker.image.tag: replaced-release
lensWorker.image.digest: sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa
- it: advertises the approved digest to standalone installers
template: backend/deployment.yaml
set:
lensWorker.image.tag: replaced-release
lensWorker.image.digest: sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LENS_WORKER_IMAGE
value: ghcr.io/berriai/litellm-lens-worker@sha256:aaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaaa
- it: refuses a malformed digest instead of falling back to the tag
template: backend/deployment.yaml
set:
lensWorker.image.digest: sha256:invalid
asserts:
- failedTemplate:
errorMessage: lensWorker.image.digest must be sha256 followed by 64 lowercase hex characters
- it: keeps the worker opt in
template: lens/deployment.yaml
asserts:
- hasDocuments:
count: 0
- it: requires a limited worker credential when enabled
template: lens/deployment.yaml
set:
lensWorker.enabled: true
asserts:
- failedTemplate:
errorMessage: lensWorker.tokenSecret.name must reference a Lens worker token
- it: uses the chart release and a secret without granting Kubernetes access
template: lens/deployment.yaml
chart:
appVersion: v1.2.3
set:
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker:v1.2.3
- equal:
path: spec.template.spec.containers[0].env[1].valueFrom.secretKeyRef
value:
name: lens-credential
key: token
- equal:
path: spec.template.spec.automountServiceAccountToken
value: false
- equal:
path: spec.template.spec.containers[0].securityContext.readOnlyRootFilesystem
value: true
- equal:
path: spec.template.spec.volumes[0].emptyDir
value:
medium: Memory
sizeLimit: 1Gi
- it: advertises the same private dev image to standalone installers
template: backend/deployment.yaml
set:
lensWorker.image.repository: registry.example/lens-worker
lensWorker.image.tag: branch-main-1234567
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LENS_WORKER_IMAGE
value: registry.example/lens-worker:branch-main-1234567
- it: supports an external gateway and a registry override
template: lens/deployment.yaml
set:
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
lensWorker.url: https://gateway.example/proxy
lensWorker.image.repository: registry.example/lens-worker
lensWorker.image.tag: branch-main-1234567
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: registry.example/lens-worker:branch-main-1234567
- equal:
path: spec.template.spec.containers[0].env[0].value
value: https://gateway.example/proxy
- it: prefixes a numeric chart release with v
template: lens/deployment.yaml
chart:
appVersion: 1.2.3-rc.4
set:
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-rc.4
- it: follows a backend image override when no worker tag is set
template: lens/deployment.yaml
set:
backend.image.tag: branch-main-1234567
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker:branch-main-1234567
- it: recommends the overridden backend release for standalone installers
template: backend/deployment.yaml
set:
backend.image.tag: v1.2.3-dev.4
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LENS_WORKER_IMAGE
value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-dev.4
- it: normalizes a numeric backend tag to the published worker tag
template: lens/deployment.yaml
set:
backend.image.tag: 1.2.3-dev.4
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-dev.4

View file

@ -629,3 +629,26 @@ ui:
affinity: {}
# Same shape as gateway.topologySpreadConstraints.
topologySpreadConstraints: []
lensWorker:
enabled: false
replicaCount: 1
image:
repository: ghcr.io/berriai/litellm-lens-worker
tag: ""
digest: ""
pullPolicy: IfNotPresent
tokenSecret:
name: ""
key: token
url: ""
tmpSizeLimit: 1Gi
resources:
requests:
cpu: 100m
memory: 256Mi
limits:
memory: 2Gi
nodeSelector: {}
tolerations: []
affinity: {}

View file

@ -39,6 +39,7 @@ fn map_error_ref(error: &Error) -> PyErr {
| Error::InsertTooLarge
| Error::ReadTooLarge => PyOverflowError::new_err(error.to_string()),
Error::InvalidRow
| Error::InvalidLimit(_)
| Error::InvalidTable
| Error::InvalidCursor(_)
| Error::AmbiguousTrace
@ -59,6 +60,7 @@ fn map_error_ref(error: &Error) -> PyErr {
Error::Cached(source) => map_error_ref(source),
Error::Storage(source) => match source {
StorageError::InvalidRow
| StorageError::InvalidLimit(_)
| StorageError::InvalidTable
| StorageError::InvalidSchema
| StorageError::EmptySql
@ -433,6 +435,13 @@ mod tests {
#[rstest]
#[case::row(Error::InvalidRow, "ValueError")]
#[case::insert_limit(Error::InvalidLimit("CLICKHOUSE_TRACE_MAX_INSERT_BYTES"), "ValueError")]
#[case::insert_timeout(
Error::Storage(litellm_storage_clickhouse::Error::InvalidLimit(
"CLICKHOUSE_INSERT_TIMEOUT_SECONDS"
)),
"ValueError"
)]
#[case::insert_budget(Error::InsertTooLarge, "OverflowError")]
#[case::scope(Error::InvalidScope, "ValueError")]
#[case::schema(Error::SchemaFailed(503), "RuntimeError")]
@ -482,6 +491,10 @@ mod tests {
#[rstest]
#[case::decode_budget(Error::Decode(litellm_traces::Error::TooLarge), "OverflowError")]
#[case::invalid_export(Error::Decode(litellm_traces::Error::InvalidPayload), "ValueError")]
#[case::invalid_decode_limit(
Error::Decode(litellm_traces::Error::InvalidLimit("OTLP_MAX_SPANS")),
"ValueError"
)]
#[case::cursor(Error::InvalidCursor("trace"), "ValueError")]
#[case::ambiguous(Error::AmbiguousTrace, "ValueError")]
#[case::changed_snapshot(Error::TraceChanged, "ValueError")]

View file

@ -2,6 +2,8 @@
pub enum Error {
#[error("invalid ClickHouse insert row")]
InvalidRow,
#[error("{0} must be a positive integer")]
InvalidLimit(&'static str),
#[error("invalid ClickHouse insert table")]
InvalidTable,
#[error("invalid ClickHouse HTTP URL")]

View file

@ -5,7 +5,19 @@ use litellm_http::Client;
use crate::{Connection, Error, valid_identifier};
const INSERT_TIMEOUT: Duration = Duration::from_secs(30);
fn insert_timeout() -> Result<Duration, Error> {
let name = "CLICKHOUSE_INSERT_TIMEOUT_SECONDS";
match std::env::var(name) {
Ok(value) => value
.parse::<u64>()
.ok()
.filter(|value| *value > 0)
.map(Duration::from_secs)
.ok_or(Error::InvalidLimit(name)),
Err(std::env::VarError::NotPresent) => Ok(Duration::from_secs(30)),
Err(_) => Err(Error::InvalidLimit(name)),
}
}
pub async fn insert_encoded_rows(
client: &Client,
@ -74,7 +86,7 @@ pub async fn insert_compressed_rows(
.append_pair("date_time_input_format", "best_effort");
let response = client
.post(url)
.timeout(INSERT_TIMEOUT)
.timeout(insert_timeout()?)
.header("Content-Encoding", "gzip")
.body(body)
.send()

View file

@ -144,3 +144,50 @@ async fn server_result_limits_allow_smaller_pages_without_retrying_other_failure
assert!(matches!(error, Error::QueryFailed(500)));
}
}
#[test]
fn insert_timeout_environment_controls_transport() {
for value in ["1", "3", "0", "invalid"] {
let result = std::process::Command::new(std::env::current_exe().unwrap())
.args(["--exact", "insert_timeout_environment_child"])
.env("LITELLM_TEST_INSERT_TIMEOUT", value)
.env("CLICKHOUSE_INSERT_TIMEOUT_SECONDS", value)
.output()
.unwrap();
assert!(
result.status.success(),
"{}",
String::from_utf8_lossy(&result.stdout)
);
}
}
#[tokio::test]
async fn insert_timeout_environment_child() {
use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method};
let Ok(value) = std::env::var("LITELLM_TEST_INSERT_TIMEOUT") else {
return;
};
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200).set_delay(std::time::Duration::from_millis(1500)))
.mount(&server)
.await;
let result = insert_encoded_rows(
&Client::no_redirect_for_test(),
&Connection::parse(&server.uri()).unwrap(),
"traces",
"otel_traces",
"token",
"{}",
)
.await;
match value.as_str() {
"1" => assert!(matches!(result, Err(Error::Transport))),
"3" => assert!(result.is_ok()),
_ => assert!(matches!(
result,
Err(Error::InvalidLimit("CLICKHOUSE_INSERT_TIMEOUT_SECONDS"))
)),
}
}

View file

@ -1,7 +1,6 @@
SELECT
request_id, response_id, trace_id, span_id, model, spend,
prompt_tokens, completion_tokens, status,
JSONExtractBool(metadata, 'synthetic_spend') AS synthetic_spend
prompt_tokens, completion_tokens, status
FROM spend_logs FINAL
WHERE start_time >= now() - INTERVAL 1 DAY
ORDER BY start_time DESC, request_id

View file

@ -2,6 +2,8 @@
pub enum Error {
#[error("invalid ClickHouse insert row")]
InvalidRow,
#[error("{0} must be a positive integer")]
InvalidLimit(&'static str),
#[error("invalid ClickHouse insert table")]
InvalidTable,
#[error("database must be a nonempty SQL identifier and retention must be positive")]

View file

@ -15,7 +15,18 @@ use time::{OffsetDateTime, format_description::well_known::Rfc3339};
use super::{Connection, Error};
use litellm_traces::Shared;
const MAX_INSERT_BYTES: usize = 64 * 1024 * 1024;
fn max_insert_bytes() -> Result<usize, Error> {
let name = "CLICKHOUSE_TRACE_MAX_INSERT_BYTES";
match std::env::var(name) {
Ok(value) => value
.parse::<usize>()
.ok()
.filter(|value| *value > 0)
.ok_or(Error::InvalidLimit(name)),
Err(std::env::VarError::NotPresent) => Ok(64 * 1024 * 1024),
Err(_) => Err(Error::InvalidLimit(name)),
}
}
pub type InsertRow = BTreeMap<String, Shared<Value>>;
@ -62,7 +73,7 @@ pub async fn insert_shared_rows(
return Ok(());
}
let received_ms = (OffsetDateTime::now_utc().unix_timestamp_nanos() / 1_000_000) as u64;
let (token, body) = prepare_insert(&rows, received_ms, MAX_INSERT_BYTES)?;
let (token, body) = prepare_insert(&rows, received_ms, max_insert_bytes()?)?;
litellm_storage_clickhouse::insert_compressed_rows(
client,
connection,

View file

@ -211,7 +211,7 @@ fn request_id(evidence: &CallEvidence) -> &str {
.flatten()
.find_map(|key| match key {
CallKey::ProviderResponse(id) => Some(id.as_str()),
CallKey::LiteLlmRequest(_) | CallKey::Transport => None,
CallKey::LiteLlmRequest(_) | CallKey::Transport | CallKey::GatewayAttempt => None,
})
.unwrap_or_default()
}

View file

@ -184,7 +184,7 @@ Duration is nanoseconds; Timestamp has nanosecond precision, spend start_time ha
{%- endblock %}
{% block missing_spend -%}
Token usage does not establish billed spend. OTLP exports without companion spend_logs rows have unknown cost; synthetic fixture spend is marked by metadata.synthetic_spend
Token usage does not establish billed spend. OTLP exports without companion spend_logs rows have unknown cost
{%- endblock %}
{% block partial_spend -%}

View file

@ -4,13 +4,11 @@ Run `cargo test -p litellm-traces-clickhouse --test queries --locked -- --test-t
Raw OTLP exports live in `crates/traces/tests/fixtures/query_*.json`. The seeded fixture decodes and normalizes them through `litellm_traces::decode_otlp` at test startup, then projects the decoded fields into ClickHouse columns. Team and key identities come from fixture setup rather than exporter claims. Root and child exports are inserted separately through the public insert API so materialized views process multiple blocks
`crates/traces/tests/fixtures/deeplite_auth_error.json` and `deeplite_swarm.json` were captured from Deeplite runs against the local proxy on 2026-10-02. The first contains a failed model call. The second contains successful model calls, searches, handoff attempts, and virtual filesystem writes. Credentials, workspace identifiers, and local user paths were redacted, and the protobuf exports were converted to OTLP JSON. Their round-trip tests check span identities, parent links, timestamps, durations, token counts, and statuses without pinning the provider's error wording
The ClickHouse round-trip test replays the `google_adk_billed_failure`, `pydantic_ai_retry`, and `deepagents_swarm` exports and checks span identities, parent links, timestamps, durations, token counts, and statuses without pinning the provider's error wording. Exported ERROR and UNSET statuses are diagnostic and do not establish a failed execution, so the test preserves incoming statuses and checks root status separately from the count of error spans, deriving both from the decoded export. Framework-specific interpretation of control-flow exceptions belongs in the instrumentation integration
The swarm capture has handoff spans marked ERROR with `ParentCommand` exception events and a root with UNSET status. These are exported diagnostic statuses, which do not establish a failed execution. The tests preserve incoming statuses and check root status separately from the count of error spans, deriving both from the decoded export. They do not infer an execution outcome from exception text, framework names, successful model calls, or output presence. Framework-specific interpretation of control-flow exceptions belongs in the instrumentation integration
For a local dashboard with linked requests and traces, run `make lens-dev ARGS=--seed` from the repository root and open `http://localhost:3000/ui/lens/`. Log in as `admin` with the master key saved in `.lens-dev/master_key`. The launcher keeps the stack running until Ctrl-C and leaves the database volumes intact
For a local dashboard with linked requests and traces, run `bash scripts/run_tracing_proxy_local.sh --seed` from the repository root and open `http://127.0.0.1:4002/ui/`. Log in as `admin` with password `sk-1234`, matching the UI E2E harness. The launcher keeps the proxy running until Ctrl-C and leaves the database volumes intact
`deeplite_swarm_spend_logs.jsonl` pairs every LLM span in the swarm export with a ClickHouse spend row. Response IDs, trace and span IDs, token counts, input, output, and timestamps come from the export. Messages and responses use the chat completion format supported by the request viewer. Spend is synthetic, set to $0.01 per request and marked in metadata, because the export does not include actual billed costs. These rows are stored here because `traces-clickhouse` owns the spend row schema
Spend rows are stored here because `traces-clickhouse` owns the spend row schema
The simple and swarm exports for all twelve SDK examples were captured on 2026-10-03 against port 4002 using `openai/gpt-6-luna`. Each export has a matching `<name>_spend_logs.jsonl` with actual proxy spend, usage, request and response IDs, messages, and timestamps. Authorization headers, provider cookies, organization and project identifiers, and local paths were redacted. OTLP identifiers and enums use their canonical JSON encodings. `metadata.fixture_capture` identifies the associated export and whether model spans contain sufficient identity to join spend

File diff suppressed because one or more lines are too long

View file

@ -138,3 +138,56 @@ fn insert_encoding_preserves_timestamp_precision_and_other_fields(
fn insert_encoding_rejects_invalid_span_timestamps(#[case] timestamp: Value) {
assert!(encode_rows(vec![BTreeMap::from([("Timestamp".into(), timestamp)])]).is_err());
}
#[test]
fn insert_byte_limit_environment_controls_transport() {
for value in ["1", "1024", "0", "invalid"] {
let result = std::process::Command::new(std::env::current_exe().unwrap())
.args(["--exact", "insert_byte_limit_environment_child"])
.env("LITELLM_TEST_INSERT_LIMIT", value)
.env("CLICKHOUSE_TRACE_MAX_INSERT_BYTES", value)
.output()
.unwrap();
assert!(
result.status.success(),
"{}",
String::from_utf8_lossy(&result.stdout)
);
}
}
#[tokio::test]
async fn insert_byte_limit_environment_child() {
let Ok(value) = std::env::var("LITELLM_TEST_INSERT_LIMIT") else {
return;
};
let server = MockServer::start().await;
Mock::given(method("POST"))
.respond_with(ResponseTemplate::new(200))
.mount(&server)
.await;
let connection = Connection::parse(&server.uri()).unwrap();
let result = insert_shared_rows(
&Client::no_redirect_for_test(),
&connection,
"traces",
InsertTable::OtelTraces,
vec![BTreeMap::from([(
"SpanId".into(),
Shared::new(json!("test")),
)])],
)
.await;
match value.as_str() {
"1" => assert!(matches!(result, Err(Error::InsertTooLarge))),
"1024" => assert!(result.is_ok()),
_ => assert!(matches!(
result,
Err(Error::InvalidLimit("CLICKHOUSE_TRACE_MAX_INSERT_BYTES"))
)),
}
assert_eq!(
server.received_requests().await.unwrap().len(),
usize::from(value == "1024")
);
}

View file

@ -244,10 +244,11 @@ async fn typed_trace_cursor_returns_the_next_fixture_trace(
}
#[rstest]
#[case::authentication_error(include_bytes!("../../traces/tests/fixtures/deeplite_auth_error.json"))]
#[case::swarm(include_bytes!("../../traces/tests/fixtures/deeplite_swarm.json"))]
#[case::billed_failure(include_bytes!("../../traces/tests/fixtures/google_adk_billed_failure.json"))]
#[case::retry(include_bytes!("../../traces/tests/fixtures/pydantic_ai_retry.json"))]
#[case::swarm(include_bytes!("../../traces/tests/fixtures/deepagents_swarm.json"))]
#[tokio::test]
async fn captured_deeplite_exports_round_trip_through_clickhouse(
async fn captured_sdk_exports_round_trip_through_clickhouse(
#[future(awt)] migrated_database: TestResult<SeededDatabase>,
admin_access: TestResult<contracts::ReadAccessParams>,
#[case] export: &[u8],

View file

@ -158,7 +158,8 @@ fn span_row(span: &DecodedSpan, team: &str, key: &str) -> BTreeMap<String, Value
.find_map(|key| match key {
litellm_traces::CallKey::LiteLlmRequest(id)
| litellm_traces::CallKey::ProviderResponse(id) => Some(id.as_str()),
litellm_traces::CallKey::Transport => None,
litellm_traces::CallKey::Transport
| litellm_traces::CallKey::GatewayAttempt => None,
})
.unwrap_or_default()
),

View file

@ -2,6 +2,8 @@
pub enum Error {
#[error("invalid OTLP trace payload")]
InvalidPayload,
#[error("{0} must be a positive integer")]
InvalidLimit(&'static str),
#[error("OTLP trace payload exceeds the decoding budget")]
TooLarge,
#[error("OTLP token count is outside the storage range")]

View file

@ -30,7 +30,7 @@ pub use normalize::{
AgentMetadata, AgentType, CallEvidence, CallEvidenceKind, CallKey, Integration, NormalizedSpan,
ObservationType,
};
pub use otlp::{DecodedEvent, DecodedSpan, decode_otlp};
pub use otlp::{DecodeLimits, DecodedEvent, DecodedSpan, decode_otlp, decode_otlp_with_limits};
pub use query::ReadQuery;
pub use query_access::QueryScope;
pub use resolve::{SpendLookup, iso_time, listed_summary, resolve_trace};

View file

@ -12,23 +12,30 @@ const SCOPES: [&str; 7] = [
];
pub(super) fn matches(context: &SpanContext<'_>) -> bool {
SCOPES.contains(&context.scope)
|| (context.scope == "litellm.gateway.client"
&& context.name == "gateway.request"
&& context
.attributes
.get("litellm.gateway.attempt")
.is_some_and(|value| value == "true")
&& context
.attributes
.get("http.request.method")
.is_some_and(|value| value == "POST"))
SCOPES.contains(&context.scope) || matches_gateway_attempt(context)
}
pub(super) fn adjust(facts: SpanFacts) -> SpanFacts {
fn matches_gateway_attempt(context: &SpanContext<'_>) -> bool {
context.scope == "litellm.gateway.client"
&& context.name == "gateway.request"
&& context
.attributes
.get("litellm.gateway.attempt")
.is_some_and(|value| value == "true")
&& context
.attributes
.get("http.request.method")
.is_some_and(|value| value == "POST")
}
pub(super) fn adjust(context: &SpanContext<'_>, facts: SpanFacts) -> SpanFacts {
SpanFacts {
role: Some(RoleEvidence::Declared(ObservationType::Framework)),
calls: CallEvidence::complete(CallKey::Transport),
calls: CallEvidence::complete(if matches_gateway_attempt(context) {
CallKey::GatewayAttempt
} else {
CallKey::Transport
}),
..facts
}
}
@ -42,7 +49,11 @@ impl Rule for HttpClient {
fn integration(&self, _: &SpanContext<'_>) -> Option<Integration> {
None
}
fn adjust(&self, _: &SpanContext<'_>, extraction: super::Extraction) -> super::Extraction {
extraction.map_facts(adjust)
fn adjust(
&self,
context: &SpanContext<'_>,
extraction: super::Extraction,
) -> super::Extraction {
extraction.map_facts(|facts| adjust(context, facts))
}
}

View file

@ -53,6 +53,7 @@ pub enum CallKey {
ProviderResponse(String),
/// The span is the HTTP request itself; LiteLLM logs its `traceparent` span id.
Transport,
GatewayAttempt,
}
impl fmt::Display for CallKey {
@ -61,6 +62,7 @@ impl fmt::Display for CallKey {
Self::LiteLlmRequest(id) => write!(formatter, "litellm_request:{id}"),
Self::ProviderResponse(id) => write!(formatter, "provider_response:{id}"),
Self::Transport => formatter.write_str("transport:"),
Self::GatewayAttempt => formatter.write_str("gateway_attempt:"),
}
}
}
@ -77,6 +79,7 @@ impl FromStr for CallKey {
Ok(Self::LiteLlmRequest(id.to_owned()))
}
Some(("transport", "")) => Ok(Self::Transport),
Some(("gateway_attempt", "")) => Ok(Self::GatewayAttempt),
_ => Err(crate::InvalidCallKey),
}
}

View file

@ -8,7 +8,7 @@ use serde::{
ser::{SerializeMap, SerializeSeq},
};
use super::limits::{Budget, MAX_ATTRIBUTES};
use super::limits::Budget;
use crate::Error;
struct AttributeWriter<'a> {
@ -33,7 +33,7 @@ pub(super) fn attributes(
values: Vec<KeyValue>,
budget: &mut Budget,
) -> Result<BTreeMap<String, String>, Error> {
if values.len() > MAX_ATTRIBUTES {
if values.len() > budget.limits.attributes {
return Err(Error::TooLarge);
}
values

View file

@ -5,14 +5,62 @@ use serde::de::{DeserializeSeed, MapAccess, SeqAccess, Visitor};
use crate::{Error, Shared};
pub(super) const MAX_DEPTH: usize = 32;
pub(super) const MAX_NODES: usize = 65_536;
pub(super) const MAX_SPANS: usize = 4_096;
pub(super) const MAX_ATTRIBUTES: usize = 256;
pub(super) const MAX_EVENTS: usize = 256;
pub(super) const MAX_DECODED_SPAN_BYTES: usize = 16 * 1024 * 1024;
#[derive(Clone, Copy, Debug)]
pub struct DecodeLimits {
pub depth: usize,
pub nodes: usize,
pub spans: usize,
pub attributes: usize,
pub events: usize,
pub links: usize,
pub decoded_span_bytes: usize,
}
pub(super) fn json_preflight(payload: &[u8]) -> Result<(), Error> {
impl Default for DecodeLimits {
fn default() -> Self {
Self {
depth: 32,
nodes: 65_536,
spans: 4_096,
attributes: 256,
events: 256,
links: 256,
decoded_span_bytes: 16 * 1024 * 1024,
}
}
}
impl DecodeLimits {
pub fn from_env() -> Result<Self, Error> {
let defaults = Self::default();
Ok(Self {
depth: env_limit("OTLP_MAX_DECODE_DEPTH", defaults.depth)?,
nodes: env_limit("OTLP_MAX_DECODE_NODES", defaults.nodes)?,
spans: env_limit("OTLP_MAX_SPANS", defaults.spans)?,
attributes: env_limit("OTLP_MAX_ATTRIBUTES", defaults.attributes)?,
events: env_limit("OTLP_MAX_EVENTS", defaults.events)?,
links: env_limit("OTLP_MAX_LINKS", defaults.links)?,
decoded_span_bytes: env_limit(
"OTLP_MAX_DECODED_SPAN_BYTES",
defaults.decoded_span_bytes,
)?,
})
}
}
fn env_limit(name: &'static str, default: usize) -> Result<usize, Error> {
match std::env::var(name) {
Ok(value) => value
.parse::<usize>()
.ok()
.filter(|value| *value > 0)
.ok_or(Error::InvalidLimit(name)),
Err(std::env::VarError::NotPresent) => Ok(default),
Err(_) => Err(Error::InvalidLimit(name)),
}
}
pub(super) fn json_preflight(payload: &[u8], limits: &DecodeLimits) -> Result<(), Error> {
let mut nodes = 0;
let mut exceeded = false;
let mut decoder = serde_json::Deserializer::from_slice(payload);
@ -20,6 +68,7 @@ pub(super) fn json_preflight(payload: &[u8]) -> Result<(), Error> {
nodes: &mut nodes,
exceeded: &mut exceeded,
depth: 0,
limits,
}
.deserialize(&mut decoder)
.and_then(|()| decoder.end());
@ -33,6 +82,7 @@ struct JsonBudget<'a> {
nodes: &'a mut usize,
exceeded: &'a mut bool,
depth: usize,
limits: &'a DecodeLimits,
}
impl<'de> DeserializeSeed<'de> for JsonBudget<'_> {
@ -40,7 +90,7 @@ impl<'de> DeserializeSeed<'de> for JsonBudget<'_> {
fn deserialize<D: serde::Deserializer<'de>>(self, decoder: D) -> Result<(), D::Error> {
*self.nodes += 1;
if *self.nodes > MAX_NODES || self.depth > MAX_DEPTH {
if *self.nodes > self.limits.nodes || self.depth > self.limits.depth {
*self.exceeded = true;
return Err(serde::de::Error::custom("OTLP structure exceeds budget"));
}
@ -79,6 +129,7 @@ impl<'de> Visitor<'de> for JsonBudget<'_> {
nodes: self.nodes,
exceeded: self.exceeded,
depth: self.depth + 1,
limits: self.limits,
})?
.is_some()
{}
@ -91,6 +142,7 @@ impl<'de> Visitor<'de> for JsonBudget<'_> {
nodes: self.nodes,
exceeded: self.exceeded,
depth: self.depth + 1,
limits: self.limits,
})?
.is_some()
{
@ -98,6 +150,7 @@ impl<'de> Visitor<'de> for JsonBudget<'_> {
nodes: self.nodes,
exceeded: self.exceeded,
depth: self.depth + 1,
limits: self.limits,
})?;
}
Ok(())
@ -146,8 +199,8 @@ impl MessageKind {
}
}
pub(super) fn protobuf_preflight(payload: &[u8]) -> Result<(), Error> {
scan_message(payload, MessageKind::Export, 0, &mut 0)
pub(super) fn protobuf_preflight(payload: &[u8], limits: &DecodeLimits) -> Result<(), Error> {
scan_message(payload, MessageKind::Export, 0, &mut 0, limits)
}
fn scan_message(
@ -155,13 +208,14 @@ fn scan_message(
kind: MessageKind,
depth: usize,
nodes: &mut usize,
limits: &DecodeLimits,
) -> Result<(), Error> {
if depth > MAX_DEPTH {
if depth > limits.depth {
return Err(Error::TooLarge);
}
while !payload.is_empty() {
*nodes += 1;
if *nodes > MAX_NODES {
if *nodes > limits.nodes {
return Err(Error::TooLarge);
}
let (tag, wire) = decode_key(&mut payload).map_err(|_| Error::InvalidPayload)?;
@ -171,7 +225,7 @@ fn scan_message(
let (message, rest) = payload
.split_at_checked(length)
.ok_or(Error::InvalidPayload)?;
scan_message(message, child, depth + 1, nodes)?;
scan_message(message, child, depth + 1, nodes, limits)?;
payload = rest;
} else {
skip_field(wire, tag, &mut payload, DecodeContext::default())
@ -183,11 +237,15 @@ fn scan_message(
pub(super) struct Budget {
remaining: usize,
pub(super) limits: DecodeLimits,
}
impl Budget {
pub(super) fn new(remaining: usize) -> Self {
Self { remaining }
pub(super) fn new(limits: DecodeLimits) -> Self {
Self {
remaining: limits.decoded_span_bytes,
limits,
}
}
pub(super) fn clone_shared<T: Clone>(

View file

@ -3,6 +3,8 @@ mod limits;
mod span;
mod wire;
pub use limits::DecodeLimits;
use serde::Serialize;
use std::collections::BTreeMap;
@ -36,6 +38,14 @@ pub struct DecodedSpan {
}
pub fn decode_otlp(body: &[u8], content_type: Option<&str>) -> Result<Vec<DecodedSpan>, Error> {
let request = wire::decode(body, content_type)?;
span::flatten(request)
decode_otlp_with_limits(body, content_type, DecodeLimits::from_env()?)
}
pub fn decode_otlp_with_limits(
body: &[u8],
content_type: Option<&str>,
limits: DecodeLimits,
) -> Result<Vec<DecodedSpan>, Error> {
let request = wire::decode(body, content_type, &limits)?;
span::flatten(request, limits)
}

View file

@ -8,15 +8,18 @@ use opentelemetry_proto::tonic::{
use super::{
DecodedEvent, DecodedSpan,
attributes::attributes,
limits::{Budget, MAX_ATTRIBUTES, MAX_DECODED_SPAN_BYTES, MAX_EVENTS, MAX_SPANS},
limits::{Budget, DecodeLimits},
};
use crate::{
Error, Shared,
normalize::{SpanContext, normalize},
};
pub(super) fn flatten(request: ExportTraceServiceRequest) -> Result<Vec<DecodedSpan>, Error> {
let mut budget = Budget::new(MAX_DECODED_SPAN_BYTES);
pub(super) fn flatten(
request: ExportTraceServiceRequest,
limits: DecodeLimits,
) -> Result<Vec<DecodedSpan>, Error> {
let mut budget = Budget::new(limits);
let mut spans = Vec::new();
for resource in request.resource_spans {
append_resource(resource, &mut budget, &mut spans)?;
@ -49,17 +52,17 @@ fn append_scope(
spans: &mut Vec<DecodedSpan>,
) -> Result<(), Error> {
let scope = scope_spans.scope.unwrap_or_default();
if scope.attributes.len() > MAX_ATTRIBUTES {
if scope.attributes.len() > budget.limits.attributes {
return Err(Error::TooLarge);
}
budget.consume(scope.name.len() + scope.version.len())?;
let scope_name: Shared<String> = scope.name.into();
let scope_version: Shared<String> = scope.version.into();
for span in scope_spans.spans {
if spans.len() >= MAX_SPANS {
if spans.len() >= budget.limits.spans {
return Err(Error::TooLarge);
}
validate_span(&span)?;
validate_span(&span, &budget.limits)?;
budget.consume(
span.name.len()
+ span.trace_state.len()
@ -85,7 +88,7 @@ fn valid_id(value: &[u8], length: usize) -> bool {
value.len() == length && value.iter().any(|byte| *byte != 0)
}
fn validate_span(span: &Span) -> Result<(), Error> {
fn validate_span(span: &Span, limits: &DecodeLimits) -> Result<(), Error> {
if !valid_id(&span.trace_id, 16)
|| !valid_id(&span.span_id, 8)
|| (!span.parent_span_id.is_empty() && !valid_id(&span.parent_span_id, 8))
@ -99,17 +102,17 @@ fn validate_span(span: &Span) -> Result<(), Error> {
{
return Err(Error::InvalidPayload);
}
if span.events.len() > MAX_EVENTS
|| span.links.len() > MAX_EVENTS
|| span.attributes.len() > MAX_ATTRIBUTES
if span.events.len() > limits.events
|| span.links.len() > limits.links
|| span.attributes.len() > limits.attributes
|| span
.links
.iter()
.any(|link| link.attributes.len() > MAX_ATTRIBUTES)
.any(|link| link.attributes.len() > limits.attributes)
|| span
.events
.iter()
.any(|event| event.attributes.len() > MAX_ATTRIBUTES)
.any(|event| event.attributes.len() > limits.attributes)
{
return Err(Error::TooLarge);
}
@ -171,7 +174,9 @@ fn decoded_span(
crate::CallKey::LiteLlmRequest(id) | crate::CallKey::ProviderResponse(id) => {
id.len() + size_of::<crate::CallKey>()
}
crate::CallKey::Transport => size_of::<crate::CallKey>(),
crate::CallKey::Transport | crate::CallKey::GatewayAttempt => {
size_of::<crate::CallKey>()
}
})
.sum::<usize>()
+ normalized.model.as_ref().map_or(0, String::len)

View file

@ -1,7 +1,7 @@
use opentelemetry_proto::tonic::collector::trace::v1::ExportTraceServiceRequest;
use prost::Message;
use super::limits::{json_preflight, protobuf_preflight};
use super::limits::{DecodeLimits, json_preflight, protobuf_preflight};
use crate::Error;
#[derive(strum::EnumString)]
@ -19,6 +19,7 @@ enum OtlpMediaType {
pub(super) fn decode(
body: &[u8],
content_type: Option<&str>,
limits: &DecodeLimits,
) -> Result<ExportTraceServiceRequest, Error> {
let media_type = content_type
.unwrap_or("application/x-protobuf")
@ -31,11 +32,11 @@ pub(super) fn decode(
let request = match media_type {
OtlpMediaType::Json => {
json_preflight(body)?;
json_preflight(body, limits)?;
serde_json::from_slice(body).map_err(|_| Error::InvalidPayload)?
}
OtlpMediaType::Protobuf => {
protobuf_preflight(body)?;
protobuf_preflight(body, limits)?;
ExportTraceServiceRequest::decode(body).map_err(|_| Error::InvalidPayload)?
}
};

View file

@ -134,11 +134,13 @@ impl<'a> Resolution<'a> {
)
}
/// The request attempts a model call made: its transport descendants, or, for bridges that
/// emit the request beside the call instead of under it, transport siblings inside the call's
/// time window when the call is the only model call under that parent.
fn transports(&self, call: usize) -> Vec<usize> {
let is_transport = |index: &usize| self.row(*index).call_keys.contains(&CallKey::Transport);
let is_transport = |index: &usize| {
self.row(*index)
.call_keys
.iter()
.any(|key| matches!(key, CallKey::Transport | CallKey::GatewayAttempt))
};
let nested: Vec<usize> = self
.graph
.descendants(call)
@ -162,7 +164,11 @@ impl<'a> Resolution<'a> {
let call_end_ns = call_start_ns + i128::from(call_row.duration_ns);
siblings
.into_iter()
.filter(is_transport)
.filter(|sibling| {
self.row(*sibling)
.call_keys
.contains(&CallKey::GatewayAttempt)
})
.filter(|sibling| {
let transport = self.row(*sibling);
let transport_start_ns = i128::from(transport.start_ns);

View file

@ -48,7 +48,9 @@ impl SpendLookup {
trace_ids: sorted(
keys()
.filter_map(|(row, key)| match key {
CallKey::Transport if !row.trace_id.is_empty() => {
CallKey::Transport | CallKey::GatewayAttempt
if !row.trace_id.is_empty() =>
{
Some(row.trace_id.clone())
}
_ => None,
@ -80,6 +82,21 @@ impl Ownership<'_> {
pub(super) type Requests<'a> = Vec<&'a SpendRow>;
#[derive(Clone, Copy, Eq, Ord, PartialEq, PartialOrd)]
enum KeyFamily {
GatewayCall,
ProviderResponse,
Transport,
}
fn key_family(key: &CallKey) -> KeyFamily {
match key {
CallKey::LiteLlmRequest(_) => KeyFamily::GatewayCall,
CallKey::ProviderResponse(_) => KeyFamily::ProviderResponse,
CallKey::Transport | CallKey::GatewayAttempt => KeyFamily::Transport,
}
}
pub(super) enum KeyMatch<'a> {
Missing,
Unique(&'a SpendRow),
@ -174,7 +191,7 @@ fn matches<'a>(
&& (spend.litellm_call_id == *id
|| (spend.litellm_call_id.is_empty() && spend.request_id == *id))
}
CallKey::Transport => {
CallKey::Transport | CallKey::GatewayAttempt => {
!row.trace_id.is_empty()
&& !row.span_id.is_empty()
&& spend.trace_id == row.trace_id
@ -216,12 +233,37 @@ pub(super) fn requests<'a>(
&& anchored
.iter()
.all(|request| request.litellm_call_id.is_empty());
let matches = keyed
let aliases: Vec<_> = keyed
.into_iter()
.filter(|(key, requests)| {
!(legacy_rows && requests.is_empty() && matches!(key, CallKey::LiteLlmRequest(_)))
})
.map(|(_, requests)| KeyMatch::new(requests))
.collect();
let families: BTreeSet<_> = aliases.iter().map(|(key, _)| key_family(key)).collect();
let compatible_rows: Vec<BTreeSet<_>> = families
.into_iter()
.map(|family| {
aliases
.iter()
.filter(|(key, _)| key_family(key) == family)
.flat_map(|(_, requests)| requests.iter().map(|request| request.identity()))
.collect()
})
.collect();
let matches = aliases
.into_iter()
.map(|(_, requests)| {
KeyMatch::new(
requests
.into_iter()
.filter(|request| {
compatible_rows
.iter()
.all(|family| family.contains(&request.identity()))
})
.collect(),
)
})
.collect();
match evidence.kind() {
CallEvidenceKind::Complete => SpendEvidence::Complete(matches),

View file

@ -172,7 +172,7 @@ fn trace_span(span: DecodedSpan) -> TraceSpansRow {
.flatten()
.find_map(|key| match key {
CallKey::ProviderResponse(id) => Some(id.clone()),
CallKey::LiteLlmRequest(_) | CallKey::Transport => None,
CallKey::LiteLlmRequest(_) | CallKey::Transport | CallKey::GatewayAttempt => None,
})
.unwrap_or_default();
TraceSpansRow {
@ -330,9 +330,7 @@ fn append_response_id(document: &mut Value, trace_id: &str, span_id: &str, respo
#[rstest]
fn captured_trace_cost_matches_spend_logs(
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")]
#[exclude("deeplite_swarm")]
spend_logs: PathBuf,
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")] spend_logs: PathBuf,
) {
let name = capture_name(&spend_logs);
let (_, capture, rows, spends) = fixture(&spend_logs);
@ -347,9 +345,7 @@ fn captured_trace_cost_matches_spend_logs(
#[rstest]
fn unrelated_sibling_transport_leaves_cost_unchanged(
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")]
#[exclude("deeplite_swarm")]
spend_logs: PathBuf,
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")] spend_logs: PathBuf,
) {
let name = capture_name(&spend_logs);
let (_, capture, rows, spends) = fixture(&spend_logs);
@ -359,7 +355,11 @@ fn unrelated_sibling_transport_leaves_cost_unchanged(
let calls: Vec<_> = rows
.iter()
.filter(|row| {
row.kind == ObservationType::Llm && !row.call_keys.contains(&CallKey::Transport)
row.kind == ObservationType::Llm
&& !row
.call_keys
.iter()
.any(|key| matches!(key, CallKey::Transport | CallKey::GatewayAttempt))
})
.cloned()
.collect();
@ -388,9 +388,7 @@ fn unrelated_sibling_transport_leaves_cost_unchanged(
#[rstest]
fn redundant_genai_response_id_keeps_call_evidence(
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")]
#[exclude("deeplite_swarm")]
spend_logs: PathBuf,
#[files("../traces-clickhouse/tests/fixtures/*_spend_logs.jsonl")] spend_logs: PathBuf,
) {
let (data, capture, _, _) = fixture(&spend_logs);
let name = capture.name;
@ -408,7 +406,9 @@ fn redundant_genai_response_id_keeps_call_evidence(
.iter()
.filter_map(|key| match key {
CallKey::ProviderResponse(id) => Some(id.clone()),
CallKey::LiteLlmRequest(_) | CallKey::Transport => None,
CallKey::LiteLlmRequest(_) | CallKey::Transport | CallKey::GatewayAttempt => {
None
}
})
.collect();
(!response_ids.is_empty()).then(|| {

File diff suppressed because one or more lines are too long

File diff suppressed because one or more lines are too long

View file

@ -746,6 +746,7 @@ fn transport_contract_keeps_independent_call_ids(span: Span) {
#[case::unrelated_scope("custom", "gateway.request", "true", "POST", false)]
#[case::unrelated_span("litellm.gateway.client", "step", "true", "POST", false)]
#[case::missing_contract("litellm.gateway.client", "gateway.request", "", "POST", false)]
#[case::disabled_contract("litellm.gateway.client", "gateway.request", "false", "POST", false)]
#[case::unrelated_method("litellm.gateway.client", "gateway.request", "true", "GET", false)]
fn gateway_attempt_contract_requires_recorded_request_boundary(
span: Span,
@ -774,7 +775,7 @@ fn gateway_attempt_contract_requires_recorded_request_boundary(
decoded.normalized.calls,
if complete {
CallEvidence::Complete(std::collections::BTreeSet::from([
CallKey::Transport,
CallKey::GatewayAttempt,
gateway,
]))
} else {

View file

@ -120,8 +120,6 @@ fn array<'a>(value: &'a Value, key: &str) -> &'a [Value] {
#[case::crewai_swarm(include_bytes!("fixtures/crewai_swarm.json"))]
#[case::deepagents_simple(include_bytes!("fixtures/deepagents_simple.json"))]
#[case::deepagents_swarm(include_bytes!("fixtures/deepagents_swarm.json"))]
#[case::deeplite_auth_error(include_bytes!("fixtures/deeplite_auth_error.json"))]
#[case::deeplite_swarm(include_bytes!("fixtures/deeplite_swarm.json"))]
#[case::google_adk_simple(include_bytes!("fixtures/google_adk_simple.json"))]
#[case::google_adk_swarm(include_bytes!("fixtures/google_adk_swarm.json"))]
#[case::langchain_simple(include_bytes!("fixtures/langchain_simple.json"))]
@ -255,6 +253,7 @@ fn llamaindex_wrapped_responses_keep_provider_call_keys(#[case] body: &[u8]) {
#[case::request(litellm_traces::CallKey::LiteLlmRequest("request:with:colons".to_owned()))]
#[case::response(litellm_traces::CallKey::ProviderResponse("response:with:colons".to_owned()))]
#[case::transport(litellm_traces::CallKey::Transport)]
#[case::gateway_attempt(litellm_traces::CallKey::GatewayAttempt)]
fn call_keys_round_trip_through_storage(#[case] key: litellm_traces::CallKey) {
assert_eq!(
key.to_string().parse::<litellm_traces::CallKey>().unwrap(),
@ -272,6 +271,8 @@ fn call_keys_round_trip_through_storage(#[case] key: litellm_traces::CallKey) {
#[case::missing_response("provider_response:")]
#[case::missing_request("litellm_request:")]
#[case::transport_id("transport:unexpected")]
#[case::gateway_attempt_separator("gateway_attempt")]
#[case::gateway_attempt_id("gateway_attempt:unexpected")]
#[case::unknown("unknown:id")]
fn malformed_call_keys_are_rejected_at_the_boundary(#[case] encoded: &str) {
assert!(encoded.parse::<litellm_traces::CallKey>().is_err());

View file

@ -1333,3 +1333,98 @@ fn resource_identity_preserves_explicit_names_and_sdk_fallbacks(
expected
);
}
#[rstest]
#[case::depth(litellm_traces::DecodeLimits { depth: 1, ..Default::default() })]
#[case::nodes(litellm_traces::DecodeLimits { nodes: 1, ..Default::default() })]
#[case::spans(litellm_traces::DecodeLimits { spans: 1, ..Default::default() })]
#[case::attributes(litellm_traces::DecodeLimits { attributes: 1, ..Default::default() })]
#[case::events(litellm_traces::DecodeLimits { events: 1, ..Default::default() })]
#[case::links(litellm_traces::DecodeLimits { links: 1, ..Default::default() })]
#[case::decoded_bytes(litellm_traces::DecodeLimits { decoded_span_bytes: 1, ..Default::default() })]
fn configurable_decode_limits_apply_to_both_wire_formats(
mut span: Span,
#[case] limits: litellm_traces::DecodeLimits,
) {
use opentelemetry_proto::tonic::{
common::v1::KeyValue,
trace::v1::span::{Event, Link},
};
use prost::Message;
span.attributes = vec![
KeyValue {
key: "a".into(),
..Default::default()
},
KeyValue {
key: "b".into(),
..Default::default()
},
];
span.events = vec![Event::default(), Event::default()];
span.links = vec![
Link {
trace_id: vec![1; 16],
span_id: vec![2; 8],
..Default::default()
};
2
];
let mut request = request_with(span.clone());
request.resource_spans[0].scope_spans[0].spans.push(span);
for (body, content_type) in [
(serde_json::to_vec(&request).unwrap(), "application/json"),
(request.encode_to_vec(), "application/x-protobuf"),
] {
assert!(matches!(
litellm_traces::decode_otlp_with_limits(&body, Some(content_type), limits),
Err(litellm_traces::Error::TooLarge)
));
assert_eq!(
litellm_traces::decode_otlp_with_limits(
&body,
Some(content_type),
litellm_traces::DecodeLimits::default()
)
.unwrap()
.len(),
2
);
}
}
#[test]
fn environment_decode_limits_are_used_and_invalid_values_fail() {
for value in ["2", "4", "0", "invalid"] {
let result = std::process::Command::new(std::env::current_exe().unwrap())
.args(["--exact", "environment_decode_limits_child"])
.env("LITELLM_TEST_DECODE_LIMIT", value)
.env("OTLP_MAX_SPANS", value)
.output()
.unwrap();
assert!(
result.status.success(),
"{}",
String::from_utf8_lossy(&result.stdout)
);
}
}
#[test]
fn environment_decode_limits_child() {
let Ok(value) = std::env::var("LITELLM_TEST_DECODE_LIMIT") else {
return;
};
let result = decode_otlp(
include_bytes!("fixtures/opentelemetry_simple.json"),
Some("application/json"),
);
match value.as_str() {
"2" => assert!(matches!(result, Err(litellm_traces::Error::TooLarge))),
"4" => assert_eq!(result.unwrap().len(), 3),
_ => assert!(matches!(
result,
Err(litellm_traces::Error::InvalidLimit("OTLP_MAX_SPANS"))
)),
}
}

View file

@ -579,7 +579,7 @@ fn sibling_transports_belong_to_the_only_model_call_under_their_parent(
10,
);
transport.trace_id = "trace".into();
transport.call_keys = vec!["transport:".parse().unwrap()];
transport.call_keys = vec![litellm_traces::CallKey::GatewayAttempt];
transport.call_evidence = Some(litellm_traces::CallEvidenceKind::Complete);
let mut rows = vec![
owned(
@ -605,11 +605,15 @@ fn sibling_transports_belong_to_the_only_model_call_under_their_parent(
}
#[rstest]
#[case::without_tool_http_sibling(None, Some(0.5))]
#[case::after_call(Some((200, 10)), Some(0.5))]
#[case::inside_call_without_spend(Some((10, 10)), None)]
#[case::without_tool_http_sibling(None, false, litellm_traces::CallKey::Transport, Some(0.5))]
#[case::after_call(Some((200, 10)), false, litellm_traces::CallKey::Transport, Some(0.5))]
#[case::inside_call_without_spend(Some((10, 10)), false, litellm_traces::CallKey::Transport, Some(0.5))]
#[case::inside_call_with_unrelated_spend(Some((10, 10)), true, litellm_traces::CallKey::Transport, Some(0.5))]
#[case::missing_gateway_attempt(Some((10, 10)), false, litellm_traces::CallKey::GatewayAttempt, None)]
fn sibling_transport_does_not_lose_model_call_spend(
#[case] transport_timing: Option<(i64, u64)>,
#[case] unrelated_spend: bool,
#[case] key: litellm_traces::CallKey,
#[case] expected: Option<f64>,
) {
let call = owned(
@ -637,24 +641,110 @@ fn sibling_transport_does_not_lose_model_call_spend(
];
let rows: Vec<_> = base_rows
.into_iter()
.chain(transport_timing.into_iter().map(|(start, duration)| {
.chain(transport_timing.map(|(start, duration)| {
let mut transport = at(
row("tool-http", "step", "GET", "framework", ""),
start,
duration,
);
transport.trace_id = "trace".into();
transport.call_keys = vec![litellm_traces::CallKey::Transport];
transport.call_keys = vec![key];
transport.call_evidence = Some(litellm_traces::CallEvidenceKind::Complete);
owned(transport, "team", "", "key")
}))
.collect();
let logged = spend("chatcmpl-1", "chatcmpl-1", "team", "", "key", 0.5);
let trace = resolve_trace("trace", "ref", &rows, &[logged]).unwrap();
let logs: Vec<_> = std::iter::once(spend("chatcmpl-1", "chatcmpl-1", "team", "", "key", 0.5))
.chain(unrelated_spend.then(|| SpendByResponseIdsRow {
trace_id: "trace".into(),
span_id: "tool-http".into(),
..spend("unrelated", "unrelated", "team", "", "key", 0.75)
}))
.collect();
let trace = resolve_trace("trace", "ref", &rows, &logs).unwrap();
assert_eq!(trace.summary.spend, expected);
assert_eq!(trace.agents[0].spend, expected);
}
#[rstest]
#[case::agreeing_ids(
litellm_traces::CallKey::Transport,
"call-a",
Some("response-a"),
Some(0.25)
)]
#[case::conflicting_gateway_id(litellm_traces::CallKey::Transport, "call-b", None, None)]
#[case::conflicting_response_id(
litellm_traces::CallKey::Transport,
"call-a",
Some("response-b"),
None
)]
#[case::conflicting_gateway_and_response(
litellm_traces::CallKey::Transport,
"call-b",
Some("response-b"),
None
)]
#[case::agreeing_gateway_attempt(
litellm_traces::CallKey::GatewayAttempt,
"call-a",
Some("response-a"),
Some(0.25)
)]
#[case::conflicting_gateway_attempt(litellm_traces::CallKey::GatewayAttempt, "call-b", None, None)]
fn gateway_attempt_identifiers_must_match_one_spend_row(
#[case] transport: litellm_traces::CallKey,
#[case] call_id: &str,
#[case] response_id: Option<&str>,
#[case] expected: Option<f64>,
) {
let keys = [
transport,
litellm_traces::CallKey::LiteLlmRequest(call_id.into()),
]
.into_iter()
.chain(response_id.map(|id| litellm_traces::CallKey::ProviderResponse(id.into())))
.collect();
let rows = [
owned(
row("agent", "", "agent", "agent", "agent"),
"team",
"",
"key",
),
owned(llm("call", "agent", "agent", ""), "team", "", "key"),
owned(
TraceSpansRow {
trace_id: "trace".into(),
call_keys: keys,
call_evidence: Some(litellm_traces::CallEvidenceKind::Complete),
..row("attempt", "call", "gateway.request", "framework", "")
},
"team",
"",
"key",
),
];
let logs = [
SpendByResponseIdsRow {
litellm_call_id: "call-a".into(),
trace_id: "trace".into(),
span_id: "attempt".into(),
..spend("request-a", "response-a", "team", "", "key", 0.25)
},
SpendByResponseIdsRow {
litellm_call_id: "call-b".into(),
trace_id: "trace".into(),
span_id: "other-attempt".into(),
..spend("request-b", "response-b", "team", "", "key", 0.5)
},
];
let trace = resolve_trace("trace", "ref", &rows, &logs).unwrap();
assert_eq!(trace.summary.spend, expected);
assert_eq!(trace.agents[0].spend, expected);
assert_eq!(trace.spans[2].spend, expected);
}
#[rstest]
#[case::legacy_row("", Some(0.5))]
#[case::other_call("other-call", None)]

View file

@ -17,6 +17,7 @@ from litellm.llms.vertex_ai.batches.transformation import (
)
from litellm.types.llms.openai import Batch
from litellm.types.utils import ModelInfo, Usage
from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS
from litellm.utils import token_counter
@ -543,6 +544,9 @@ def _extract_file_access_credentials(litellm_params: dict | None) -> dict:
"max_retries",
"_litellm_internal_model_credentials",
*AWS_CREDENTIAL_KWARGS_KEYS,
# A federated deployment holds no api_key, so without these the fetch that reads a
# finished batch's output has nothing to authenticate with and its cost is never billed.
*sorted(ANTHROPIC_WIF_KWARGS_KEYS),
)
for key in credential_keys:
if key in litellm_params:

View file

@ -200,7 +200,7 @@ def create_batch(
LiteLLM Equivalent of POST: https://api.openai.com/v1/batches
"""
try:
optional_params: Final = GenericLiteLLMParams(**kwargs)
optional_params: Final = GenericLiteLLMParams.model_validate(kwargs)
litellm_call_id: Final = kwargs.get("litellm_call_id", None)
proxy_server_request: Final = kwargs.get("proxy_server_request", None)
model_info: Final = kwargs.get("model_info", None)
@ -217,7 +217,7 @@ def create_batch(
)
_is_async: Final = kwargs.pop("acreate_batch", False) is True
litellm_params: Final = dict(GenericLiteLLMParams(**kwargs))
litellm_params: Final = dict(GenericLiteLLMParams.model_validate(kwargs))
litellm_logging_obj: Final[LiteLLMLoggingObj] = cast(LiteLLMLoggingObj, kwargs.get("litellm_logging_obj", None))
### TIMEOUT LOGIC ###
timeout: Final = _resolve_timeout(optional_params, kwargs, custom_llm_provider)
@ -530,6 +530,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
)
api_key = optional_params.api_key or litellm.api_key or litellm.azure_key or get_secret_str("ANTHROPIC_API_KEY")
batch_params: Final = dict(litellm_params)
response = anthropic_batches_instance.retrieve_batch(
_is_async=_is_async,
batch_id=batch_id,
@ -537,6 +538,7 @@ def _handle_retrieve_batch_providers_without_provider_config(
api_key=api_key,
timeout=timeout,
max_retries=optional_params.max_retries,
litellm_params=batch_params,
)
else:
raise litellm.exceptions.BadRequestError(
@ -573,7 +575,7 @@ def retrieve_batch(
LiteLLM Equivalent of GET https://api.openai.com/v1/batches/{batch_id}
"""
try:
optional_params: Final = GenericLiteLLMParams(**kwargs)
optional_params: Final = GenericLiteLLMParams.model_validate(kwargs)
litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None)
### TIMEOUT LOGIC ###
timeout = optional_params.timeout or kwargs.get("request_timeout", 600) or 600
@ -755,7 +757,7 @@ def list_batches(
"""
try:
# set API KEY
optional_params: Final = GenericLiteLLMParams(**kwargs)
optional_params: Final = GenericLiteLLMParams.model_validate(kwargs)
litellm_params: Final = get_litellm_params(
custom_llm_provider=custom_llm_provider,
**kwargs,
@ -956,7 +958,7 @@ def cancel_batch(
verbose_logger.exception(
"litellm.batches.main.py::cancel_batch() - Error inferring custom_llm_provider - %s", e
)
optional_params: Final = GenericLiteLLMParams(**kwargs)
optional_params: Final = GenericLiteLLMParams.model_validate(kwargs)
litellm_params: Final = get_litellm_params(
custom_llm_provider=custom_llm_provider,
**kwargs,

View file

@ -28,6 +28,7 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
from litellm.llms.anthropic.common_utils import (
is_claude_code_one_shot_subagent_request,
supports_anthropic_cache_control,
tool_call_is_rebuilt_as_server_tool_use,
)
from litellm.types.integrations.anthropic_cache_control_hook import (
GATEWAY_INJECTED_CACHE_METADATA_KEY,
@ -122,7 +123,30 @@ def targets_openai_api(api_base: object) -> bool:
def _carries_cache_breakpoint(block: object) -> bool:
return isinstance(block, dict) and any(block.get(key) is not None for key in CACHE_BREAKPOINT_KEYS)
return any(_attribute_or_key(block, key) is not None for key in CACHE_BREAKPOINT_KEYS)
def _attribute_or_key(value: object, key: str) -> object | None:
if hasattr(value, key):
return getattr(value, key)
if isinstance(value, Mapping):
return value.get(key)
return None
def _as_object_list(value: object | None) -> list[object] | None:
if not isinstance(value, list):
return None
return _validated_object_list(value)
def _tool_call_carries_cache_breakpoint(tool_call: object, message: object) -> bool:
if _attribute_or_key(tool_call, "cache_control") is None:
return False
return not tool_call_is_rebuilt_as_server_tool_use(
_attribute_or_key(tool_call, "id"), _attribute_or_key(message, "provider_specific_fields")
)
def _tool_carries_cache_breakpoint(tool: object) -> bool:
@ -471,13 +495,16 @@ class AnthropicCacheControlHook(CustomPromptManagement):
@staticmethod
def _count_cache_control_blocks(message: object) -> int:
if not isinstance(message, dict):
return 0
count = 1 if _carries_cache_breakpoint(message) else 0
content: Final = message.get("content")
if isinstance(content, list):
count += sum(1 for block in content if _carries_cache_breakpoint(block))
return count
message_count: Final = 1 if _carries_cache_breakpoint(message) else 0
content: Final = _as_object_list(_attribute_or_key(message, "content"))
content_count: Final = sum(1 for block in content if _carries_cache_breakpoint(block)) if content else 0
tool_calls: Final = _as_object_list(_attribute_or_key(message, "tool_calls"))
tool_call_count: Final = (
sum(1 for tool_call in tool_calls if _tool_call_carries_cache_breakpoint(tool_call, message))
if tool_calls
else 0
)
return message_count + content_count + tool_call_count
@staticmethod
def _message_has_cache_control(message: AllMessageValues) -> bool:

View file

@ -11,6 +11,7 @@ from litellm.litellm_core_utils.core_helpers import normalize_drop_params
from litellm.llms.openai.data_residency import infer_openai_data_residency
from litellm.types.litellm_params import MAX_CONTROL_INT_DIGITS, ControlOptions
from litellm.types.router import CustomPricingLiteLLMParams
from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS, OPENAI_WIF_KWARGS_KEYS
AWS_CREDENTIAL_KWARGS_KEYS: Final = frozenset(
{
@ -70,6 +71,8 @@ OPTIONAL_KWARGS_KEYS: Final = (
}
)
| AWS_CREDENTIAL_KWARGS_KEYS
| ANTHROPIC_WIF_KWARGS_KEYS
| OPENAI_WIF_KWARGS_KEYS
| frozenset(CustomPricingLiteLLMParams.model_fields)
)

View file

@ -24,6 +24,32 @@ IMAGE_EDIT_HEALTH_CHECK_PROMPT: Final = (
"Add a small yellow star in the top right corner of this simple drawing of a blue circle on a white background"
)
ANTHROPIC_MESSAGES_HEALTH_CHECK_MAX_TOKENS: Final = 16
def native_health_check_mode(model: str, custom_llm_provider: str | None) -> Literal["anthropic_messages"] | None:
if custom_llm_provider != "bedrock_mantle":
return None
from litellm.llms.bedrock_mantle.common_utils import mantle_health_check_mode
return mantle_health_check_mode(model)
def _cost_map_mode(model: str) -> str | None:
import litellm
from litellm.litellm_core_utils.health_check_utils import OPTIONAL_STR
return OPTIONAL_STR.validate_python(litellm.model_cost.get(model, {}).get("mode"))
def default_health_check_mode(requested_model: str, model: str, custom_llm_provider: str) -> str:
return (
native_health_check_mode(model=model, custom_llm_provider=custom_llm_provider)
or _cost_map_mode(requested_model)
or _cost_map_mode(model)
or "chat"
)
def get_image_file_for_health_check() -> bytes:
"""Return the image used for health checks."""
@ -167,6 +193,7 @@ class HealthCheckHelpers:
"realtime",
"batch",
"responses",
"anthropic_messages",
"ocr",
"evaluation",
],
@ -254,6 +281,13 @@ class HealthCheckHelpers:
**_filter_model_params(model_params=model_params),
input=prompt or "test",
),
"anthropic_messages": lambda: litellm.anthropic_messages(
**{
"max_tokens": ANTHROPIC_MESSAGES_HEALTH_CHECK_MAX_TOKENS,
"messages": [{"role": "user", "content": prompt or "test"}],
**model_params,
}
),
"ocr": lambda: litellm.aocr(
**_filter_model_params(model_params=model_params),
document=_ocr_health_check_document(model=model, custom_llm_provider=custom_llm_provider),

View file

@ -9,6 +9,7 @@ from pydantic import TypeAdapter
from litellm.types.decisions import DecisionsCallParams
DECISIONS_CALL_PARAMS: Final[TypeAdapter[DecisionsCallParams]] = TypeAdapter(DecisionsCallParams)
OPTIONAL_STR: Final[TypeAdapter[str | None]] = TypeAdapter(str | None)
def _filter_model_params(model_params: dict) -> dict:

View file

@ -1729,11 +1729,16 @@ def convert_function_to_anthropic_tool_invoke(
raise e
def _find_server_tool_result(
ANTHROPIC_SERVER_TOOL_USE_ID_PREFIX: Final = "srvtoolu_"
def find_anthropic_server_tool_result(
tool_id: str,
web_search_results: Sequence[object] | None,
tool_results: Sequence[object] | None,
) -> dict[str, object] | None:
if not tool_id.startswith(ANTHROPIC_SERVER_TOOL_USE_ID_PREFIX):
return None
candidates: Final = (*(web_search_results or ()), *(tool_results or ()))
return next(
(result for result in candidates if isinstance(result, dict) and result.get("tool_use_id") == tool_id),
@ -1808,11 +1813,7 @@ def convert_to_anthropic_tool_invoke(
context="Anthropic tool invoke",
)
server_tool_result = (
_find_server_tool_result(tool_id, web_search_results, tool_results)
if tool_id.startswith("srvtoolu_")
else None
)
server_tool_result = find_anthropic_server_tool_result(tool_id, web_search_results, tool_results)
if server_tool_result is not None:
anthropic_tool_invoke.append(
{

View file

@ -42,6 +42,7 @@ class AnthropicBatchesHandler:
timeout: float | httpx.Timeout,
max_retries: int | None,
logging_obj: LiteLLMLoggingObj | None = None,
litellm_params: dict | None = None, # mutable-ok: handed straight to validate_environment
) -> LiteLLMBatch:
"""
Async: Retrieve a batch from Anthropic.
@ -60,9 +61,7 @@ class AnthropicBatchesHandler:
# Resolve API credentials
api_base = api_base or self.anthropic_model_info.get_api_base(api_base)
api_key = api_key or self.anthropic_model_info.get_api_key()
if not api_key:
raise ValueError("Missing Anthropic API Key")
resolved_litellm_params: Final = litellm_params if litellm_params is not None else {}
# Create a minimal logging object if not provided
if logging_obj is None:
@ -85,16 +84,18 @@ class AnthropicBatchesHandler:
api_base=api_base,
batch_id=batch_id,
optional_params={},
litellm_params={},
litellm_params=resolved_litellm_params,
)
# Validate environment and get headers
headers: Final = self.provider_config.validate_environment(
# Validate environment and get headers. Offloaded to a worker thread: a WIF token
# exchange here would otherwise block the event loop.
headers: Final = await asyncio.to_thread(
self.provider_config.validate_environment,
headers={},
model="",
messages=[],
optional_params={},
litellm_params={},
litellm_params=resolved_litellm_params,
api_key=api_key,
api_base=api_base,
)
@ -130,6 +131,7 @@ class AnthropicBatchesHandler:
timeout: float | httpx.Timeout,
max_retries: int | None,
logging_obj: LiteLLMLoggingObj | None = None,
litellm_params: dict | None = None, # mutable-ok: handed straight to validate_environment
) -> LiteLLMBatch | Coroutine[Any, Any, LiteLLMBatch]:
"""
Retrieve a batch from Anthropic.
@ -154,6 +156,7 @@ class AnthropicBatchesHandler:
timeout=timeout,
max_retries=max_retries,
logging_obj=logging_obj,
litellm_params=litellm_params,
)
else:
return asyncio.run(
@ -164,5 +167,6 @@ class AnthropicBatchesHandler:
timeout=timeout,
max_retries=max_retries,
logging_obj=logging_obj,
litellm_params=litellm_params,
)
)

View file

@ -13,6 +13,8 @@ from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.types.llms.openai import AllMessageValues, CreateBatchRequest
from litellm.types.utils import LiteLLMBatch, LlmProviders, ModelResponse
from ..common_utils import merge_anthropic_beta_headers, without_caller_credential_headers
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer
@ -69,24 +71,30 @@ class AnthropicBatchesConfig(BaseBatchesConfig):
api_base: str | None = None,
) -> dict:
"""Validate and prepare environment-specific headers and parameters."""
if api_base is None and isinstance(litellm_params, dict):
api_base = litellm_params.get("api_base")
auth_header: Final = self.anthropic_model_info.get_auth_header(api_key, api_base)
params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None
if api_base is None and params_mapping is not None:
api_base = params_mapping.get("api_base")
auth_header: Final = self.anthropic_model_info.get_auth_header(
api_key, api_base, litellm_params=params_mapping, allow_workload_identity=True
)
if auth_header is None:
raise ValueError(
"Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params"
)
_headers: Final = {
merged_beta: Final = merge_anthropic_beta_headers(
merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")),
"message-batches-2024-09-24",
)
# The deployment's own credential is applied below, so a caller-supplied one must not
# ride along: without this a minted federation Bearer travels beside the caller's x-api-key.
return {
**without_caller_credential_headers(headers),
"accept": "application/json",
"anthropic-version": "2023-06-01",
"content-type": "application/json",
**auth_header,
"anthropic-beta": merged_beta,
}
_headers.update(auth_header)
# Add beta header for message batches
if "anthropic-beta" not in headers:
headers["anthropic-beta"] = "message-batches-2024-09-24"
headers.update(_headers)
return headers
def get_complete_batch_url(
self,

View file

@ -3,7 +3,7 @@ import re
import time
from collections.abc import Callable, Mapping, Sequence
from types import MappingProxyType
from typing import TYPE_CHECKING, Any, Final, NoReturn, cast
from typing import TYPE_CHECKING, Any, ClassVar, Final, NoReturn, cast
import httpx
from pydantic import BaseModel, ValidationError
@ -296,6 +296,8 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig):
to pass metadata to anthropic, it's {"user_id": "any-relevant-information"}
"""
_workload_identity_eligible: ClassVar[bool] = True
max_tokens: int | None = None
stop_sequences: list | None = None
temperature: int | None = None

View file

@ -7,10 +7,11 @@ import re
from collections.abc import Mapping, MutableMapping, Sequence
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Any, Final, Literal, TypeVar
from typing import Any, ClassVar, Final, Literal, TypeVar
from urllib.parse import quote
import httpx
from pydantic import BaseModel, ConfigDict, StrictBool, TypeAdapter, ValidationError
from pydantic import BaseModel, ConfigDict, Field, StrictBool, TypeAdapter, ValidationError
import litellm
from litellm.constants import (
@ -26,10 +27,18 @@ from litellm.litellm_core_utils.prompt_templates.common_utils import (
)
from litellm.litellm_core_utils.prompt_templates.factory import (
THOUGHT_SIGNATURE_SEPARATOR,
find_anthropic_server_tool_result,
)
from litellm.litellm_core_utils.prompt_templates.mid_conversation_system import message_field, parts_of
from litellm.llms.anthropic.wif import (
aget_anthropic_wif_token,
anthropic_base_without_chat_suffix,
get_anthropic_wif_token,
warn_if_static_credential_shadows_federation,
)
from litellm.llms.base_llm.base_utils import BaseLLMModelInfo, BaseTokenCounter
from litellm.llms.base_llm.chat.transformation import BaseLLMException
from litellm.proxy._types import SpecialHeaders
from litellm.types.llms.anthropic import (
ANTHROPIC_HOSTED_TOOLS,
ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER,
@ -222,6 +231,29 @@ def _strip_bedrock_id_suffixes(model: str) -> str:
)
_SERVER_OWNED_AUTH_HEADERS: Final = SpecialHeaders.litellm_credential_header_names()
_WIF_ELIGIBILITY_ATTR: Final = "_workload_identity_eligible"
def without_caller_credential_headers(headers: Mapping[str, str]) -> Mapping[str, str]:
"""``headers`` minus every header that authenticates the caller to litellm.
The deployment's own credential is applied on top of the result, so a caller-supplied
credential must not survive into the upstream request: without this a minted federation
Bearer travels beside the caller's own ``x-api-key``, and Anthropic sees two credentials.
"""
return MappingProxyType(
{name: value for name, value in headers.items() if name.lower() not in _SERVER_OWNED_AUTH_HEADERS}
)
def config_allows_workload_identity(config: object) -> bool:
"""A federation token is an Anthropic-org credential and its exchange POSTs the workload's OIDC
assertion to the deployment's own host, so eligibility is declared per class and read from that
class's own ``__dict__``: a subclass written for another provider inherits nothing."""
return type(config).__dict__.get(_WIF_ELIGIBILITY_ATTR, False) is True
def is_anthropic_oauth_key(value: str | None) -> bool:
"""Check if a value contains an Anthropic OAuth token (sk-ant-oat*)."""
if value is None:
@ -240,12 +272,22 @@ def resolve_used_client_oauth_token(client_sent_oauth_token: object, custom_llm_
return client_sent_oauth_token and custom_llm_provider in ANTHROPIC_OAUTH_FORWARD_PROVIDERS
def _merge_beta_headers(existing: str | None, new_beta: str) -> str:
"""Merge a new beta value into an existing comma-separated anthropic-beta header."""
if not existing:
return new_beta
betas: Final = {b.strip() for b in existing.split(",") if b.strip()}
betas.add(new_beta)
def _beta_header_values(side: str | Sequence[str] | None) -> tuple[str, ...]:
if not side:
return ()
if isinstance(side, str):
return (side,)
return tuple(entry for entry in side if isinstance(entry, str))
def merge_anthropic_beta_headers(existing: str | Sequence[str] | None, new_beta: str | Sequence[str] | None) -> str:
"""Merge anthropic-beta header values, deduplicated and sorted.
Either side may arrive as a list rather than a comma-separated string: the Skills surface
accepted a list-valued header before it shared this helper, and callers still send one.
"""
joined: Final = ",".join(_beta_header_values(existing) + _beta_header_values(new_beta))
betas: Final = frozenset(b.strip() for b in joined.split(",") if b.strip())
return ",".join(sorted(betas))
@ -272,7 +314,9 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup
):
headers.pop(name)
headers["authorization"] = auth_header
headers["anthropic-beta"] = _merge_beta_headers(headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER)
headers["anthropic-beta"] = merge_anthropic_beta_headers(
headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER
)
headers["anthropic-dangerous-direct-browser-access"] = "true"
return headers, api_key
# Check api_key directly (standard chat/completion flow)
@ -280,7 +324,9 @@ def optionally_handle_anthropic_oauth(headers: dict, api_key: str | None) -> tup
for name in tuple(header_name for header_name in headers if header_name.lower() == "x-api-key"):
headers.pop(name)
headers["authorization"] = f"Bearer {api_key}"
headers["anthropic-beta"] = _merge_beta_headers(headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER)
headers["anthropic-beta"] = merge_anthropic_beta_headers(
headers.get("anthropic-beta"), ANTHROPIC_OAUTH_BETA_HEADER
)
headers["anthropic-dangerous-direct-browser-access"] = "true"
return headers, api_key
@ -316,7 +362,79 @@ class AnthropicError(BaseLLMException):
super().__init__(status_code=status_code, message=message, headers=headers)
_MODEL_LIST_PAGE_CAP: Final = 20
def _litellm_params_str(litellm_params: Mapping[str, object] | None, key: str) -> str | None:
value: Final = litellm_params.get(key) if litellm_params is not None else None
return value if isinstance(value, str) else None
class _AnthropicModelListEntry(BaseModel):
id: str
class _AnthropicModelsPage(BaseModel):
data: Sequence[_AnthropicModelListEntry] = Field(default_factory=tuple)
has_more: bool = False
last_id: str | None = None
def _sanitized_anthropic_error(response: httpx.Response, detail: str | None = None) -> str:
"""A provider error detail built only from structured fields, never ``response.text``
verbatim: the raw body is untrusted content the caller of ``/v1/models`` did not ask for
and should not have echoed back to it wholesale."""
if detail is not None:
return f"HTTP {response.status_code}: {detail}"
try:
body: Final = response.json()
except ValueError:
return f"HTTP {response.status_code}"
error: Final = body.get("error") if isinstance(body, dict) else None
message: Final = error.get("message") if isinstance(error, dict) else None
return f"HTTP {response.status_code}: {message}" if isinstance(message, str) else f"HTTP {response.status_code}"
def _fetch_anthropic_models_page(
api_base: str, headers: Mapping[str, str], after_id: str | None
) -> _AnthropicModelsPage:
# after_id rides the URL because the client mutates the params mapping it is handed,
# which a read-only one cannot support
query: Final = f"?after_id={quote(after_id)}" if after_id else ""
response: Final = litellm.module_level_client.get(
url=f"{api_base}/v1/models{query}",
headers=headers,
follow_redirects=False,
)
try:
response.raise_for_status()
except httpx.HTTPStatusError:
raise Exception(f"Failed to fetch models from Anthropic. {_sanitized_anthropic_error(response)}") from None
try:
return _AnthropicModelsPage.model_validate(response.json())
except ValueError as e:
raise Exception(
f"Failed to fetch models from Anthropic. {_sanitized_anthropic_error(response, detail=str(e))}"
) from None
def _fetch_anthropic_model_ids(
api_base: str, headers: Mapping[str, str], after_id: str | None, pages_left: int
) -> tuple[str, ...]:
collected: tuple[str, ...] = () # rebind-ok: accumulates one page of ids per iteration
cursor: str | None = after_id # rebind-ok: advances to each page's last_id
for _ in range(max(pages_left, 0)):
page = _fetch_anthropic_models_page(api_base, headers, cursor)
collected += tuple(entry.id for entry in page.data)
if not page.has_more or page.last_id is None:
return collected
cursor = page.last_id
raise Exception(f"Anthropic /v1/models did not terminate within {_MODEL_LIST_PAGE_CAP} pages.")
class AnthropicModelInfo(BaseLLMModelInfo):
_workload_identity_eligible: ClassVar[bool] = True
def is_cache_control_set(self, messages: list[AllMessageValues]) -> bool:
"""
Return if {"cache_control": ..} in message content block
@ -940,7 +1058,9 @@ class AnthropicModelInfo(BaseLLMModelInfo):
return list(set(betas).union(thinking_display_betas, tool_change_betas))
@staticmethod
def _make_api_key_auth_header(api_key: str, api_base: str | None, use_bearer_for_custom_base: bool = False) -> dict:
def _make_api_key_auth_header(
api_key: str, api_base: str | None, use_bearer_for_custom_base: bool = False
) -> Mapping[str, str]:
if use_bearer_for_custom_base and (
api_base and "api.anthropic.com" not in api_base and not api_key.startswith("sk-ant-")
):
@ -948,6 +1068,33 @@ class AnthropicModelInfo(BaseLLMModelInfo):
return {"authorization": value}
return {"x-api-key": api_key}
def _credential_headers(
self,
*,
api_key: str | None,
auth_token: str | None,
api_base: str | None,
use_bearer_for_custom_base: bool,
wif_minted: bool,
betas: set[str], # mutable-ok: the caller's beta accumulator, appended to by the oauth tier
) -> Mapping[str, str]:
"""The credential tier walk: a consumer OAuth token, then ANTHROPIC_AUTH_TOKEN, then an api key.
A server-minted federation token takes the same Bearer shape as a consumer OAuth token but is
not browser-forwarded, so it does not get the direct-browser-access header.
"""
if api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX):
betas.add(ANTHROPIC_OAUTH_BETA_HEADER)
oauth_headers: Final = {"authorization": f"Bearer {api_key}"}
if wif_minted:
return oauth_headers
return {**oauth_headers, "anthropic-dangerous-direct-browser-access": "true"}
if auth_token and not api_key:
return {"authorization": f"Bearer {auth_token}"}
if api_key:
return self._make_api_key_auth_header(api_key, api_base, use_bearer_for_custom_base)
return {}
def get_anthropic_headers(
self,
api_key: str | None = None,
@ -972,6 +1119,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
is_mid_conversation_output_config_used: bool = False,
is_thinking_display_updates_used: bool = False,
is_mid_conversation_tool_change_used: bool = False,
wif_minted: bool = False,
) -> dict:
betas: Final = set()
# Anthropic no longer requires the prompt-caching beta header
@ -1010,20 +1158,21 @@ class AnthropicModelInfo(BaseLLMModelInfo):
if is_mid_conversation_output_config_used:
betas.add(ANTHROPIC_MID_CONVERSATION_OUTPUT_CONFIG_BETA_HEADER)
_is_oauth: Final = api_key and api_key.startswith(ANTHROPIC_OAUTH_TOKEN_PREFIX)
headers: Final = {
"anthropic-version": anthropic_version or "2023-06-01",
"accept": "application/json",
"content-type": "application/json",
}
if _is_oauth:
headers["authorization"] = f"Bearer {api_key}"
headers["anthropic-dangerous-direct-browser-access"] = "true"
betas.add(ANTHROPIC_OAUTH_BETA_HEADER)
elif auth_token and not api_key:
headers["authorization"] = f"Bearer {auth_token}"
elif api_key:
headers.update(self._make_api_key_auth_header(api_key, api_base, use_bearer_for_custom_base))
headers.update(
self._credential_headers(
api_key=api_key,
auth_token=auth_token,
api_base=api_base,
use_bearer_for_custom_base=use_bearer_for_custom_base,
wif_minted=wif_minted,
betas=betas,
)
)
if user_anthropic_beta_headers is not None:
betas.update(user_anthropic_beta_headers)
@ -1055,10 +1204,11 @@ class AnthropicModelInfo(BaseLLMModelInfo):
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
if api_base is None and isinstance(litellm_params, dict):
api_base = litellm_params.get("api_base")
params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None
if api_base is None and params_mapping is not None:
api_base = params_mapping.get("api_base")
use_bearer_for_custom_base: Final[bool] = bool(
isinstance(litellm_params, dict) and litellm_params.get("use_bearer_for_custom_base", False)
params_mapping is not None and params_mapping.get("use_bearer_for_custom_base", False)
)
# Check for Anthropic OAuth token in headers
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
@ -1067,9 +1217,25 @@ class AnthropicModelInfo(BaseLLMModelInfo):
auth_token: str | None = None
if api_key is None:
auth_token = AnthropicModelInfo.get_auth_token()
if api_key is None and auth_token is None:
if (api_key is not None or auth_token is not None) and config_allows_workload_identity(self):
warn_if_static_credential_shadows_federation(params_mapping, model)
wif_token: Final = (
get_anthropic_wif_token(params_mapping, api_base, model)
if api_key is None and auth_token is None and config_allows_workload_identity(self)
else None
)
wif_minted: Final = wif_token is not None
resolved_api_key: Final = wif_token if wif_token is not None else api_key
if resolved_api_key is None and auth_token is None:
raise litellm.AuthenticationError(
message="Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` in your environment vars",
message=(
"Missing Anthropic API Key - A call is being made to anthropic but no key is set either in the "
"environment variables or via params. Please set `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` "
"in your environment vars, or configure workload identity federation via "
"`ANTHROPIC_FEDERATION_RULE_ID`, `ANTHROPIC_ORGANIZATION_ID`, "
"`ANTHROPIC_SERVICE_ACCOUNT_ID` and "
"`ANTHROPIC_IDENTITY_TOKEN_FILE` (or `ANTHROPIC_IDENTITY_TOKEN`)"
),
llm_provider="anthropic",
model=model,
)
@ -1095,7 +1261,7 @@ class AnthropicModelInfo(BaseLLMModelInfo):
computer_tool_used=computer_tool_used,
prompt_caching_set=prompt_caching_set,
pdf_used=pdf_used,
api_key=api_key,
api_key=resolved_api_key,
auth_token=auth_token,
file_id_used=file_id_used,
is_mid_conversation_output_config_used=is_mid_conversation_output_config_used,
@ -1113,11 +1279,12 @@ class AnthropicModelInfo(BaseLLMModelInfo):
container_with_skills_used=container_with_skills_used,
api_base=api_base,
use_bearer_for_custom_base=use_bearer_for_custom_base,
wif_minted=wif_minted,
)
headers = {**headers, **anthropic_headers}
caller_headers: Final = without_caller_credential_headers(headers) if wif_minted else headers
return headers
return {**caller_headers, **anthropic_headers}
@staticmethod
def get_api_base(api_base: str | None = None) -> str | None:
@ -1132,9 +1299,13 @@ class AnthropicModelInfo(BaseLLMModelInfo):
@staticmethod
def get_api_key(api_key: str | None = None) -> str | None:
from litellm.secret_managers.main import get_secret_str
"""An empty or whitespace-only key counts as unset: it can never authenticate anything, and
treating it as set would silently outrank workload identity federation."""
from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str
return api_key or get_secret_str("ANTHROPIC_API_KEY")
return normalize_nonempty_secret_str(api_key) or normalize_nonempty_secret_str(
get_secret_str("ANTHROPIC_API_KEY")
)
@staticmethod
def get_auth_token(auth_token: str | None = None) -> str | None:
@ -1143,61 +1314,130 @@ class AnthropicModelInfo(BaseLLMModelInfo):
Unlike api_key (which uses X-Api-Key header), auth_token uses
Authorization: Bearer header, matching the official Anthropic SDK behavior.
"""
from litellm.secret_managers.main import get_secret_str
from litellm.secret_managers.main import get_secret_str, normalize_nonempty_secret_str
return auth_token or get_secret_str("ANTHROPIC_AUTH_TOKEN")
return normalize_nonempty_secret_str(auth_token) or normalize_nonempty_secret_str(
get_secret_str("ANTHROPIC_AUTH_TOKEN")
)
@staticmethod
def get_auth_header(
api_key: str | None = None,
api_base: str | None = None,
use_bearer_for_custom_base: bool = False,
) -> dict | None:
litellm_params: Mapping[str, object] | None = None,
allow_workload_identity: bool = False,
) -> Mapping[str, str] | None:
"""Resolve Anthropic credentials and return the appropriate auth header dict.
Checks ANTHROPIC_API_KEY first (-> x-api-key or Bearer depending on
use_bearer_for_custom_base), then ANTHROPIC_AUTH_TOKEN (-> Authorization: Bearer).
Returns None if neither is available.
use_bearer_for_custom_base), then ANTHROPIC_AUTH_TOKEN (-> Authorization: Bearer),
then workload identity federation (-> Authorization: Bearer with a minted
sk-ant-oat01 token, honoring anthropic_* litellm_params when provided). Every
Bearer built from an sk-ant-oat token carries the mandatory oauth anthropic-beta.
Returns None if no credential source is available.
"""
static_header: Final = AnthropicModelInfo._static_auth_header(api_key, api_base, use_bearer_for_custom_base)
if static_header is not None:
return static_header
if not allow_workload_identity:
return None
wif_token: Final = get_anthropic_wif_token(litellm_params, api_base, "")
if wif_token is not None:
return AnthropicModelInfo._oauth_bearer_header(wif_token)
return None
@staticmethod
async def aget_auth_header(
api_key: str | None = None,
api_base: str | None = None,
use_bearer_for_custom_base: bool = False,
litellm_params: Mapping[str, object] | None = None,
allow_workload_identity: bool = False,
) -> Mapping[str, str] | None:
"""Async counterpart of get_auth_header: the WIF tier can block on a token
exchange POST, so async callers await it off the event loop."""
static_header: Final = AnthropicModelInfo._static_auth_header(api_key, api_base, use_bearer_for_custom_base)
if static_header is not None:
return static_header
if not allow_workload_identity:
return None
wif_token: Final = await aget_anthropic_wif_token(litellm_params, api_base, "")
if wif_token is not None:
return AnthropicModelInfo._oauth_bearer_header(wif_token)
return None
@staticmethod
def _static_auth_header(
api_key: str | None,
api_base: str | None,
use_bearer_for_custom_base: bool,
) -> Mapping[str, str] | None:
resolved_key: Final = AnthropicModelInfo.get_api_key(api_key)
if resolved_key is not None:
if is_anthropic_oauth_key(resolved_key):
return {"authorization": f"Bearer {resolved_key}"}
return AnthropicModelInfo._oauth_bearer_header(resolved_key)
return AnthropicModelInfo._make_api_key_auth_header(resolved_key, api_base, use_bearer_for_custom_base)
auth_token: Final = AnthropicModelInfo.get_auth_token()
if auth_token is not None:
return {"authorization": f"Bearer {auth_token}"}
return None
@staticmethod
def _oauth_bearer_header(token: str) -> Mapping[str, str]:
return {"authorization": f"Bearer {token}", "anthropic-beta": ANTHROPIC_OAUTH_BETA_HEADER}
@staticmethod
def get_base_model(model: str | None = None) -> str | None:
return model.replace("anthropic/", "") if model else None
def get_models(self, api_key: str | None = None, api_base: str | None = None) -> list[str]:
api_base = AnthropicModelInfo.get_api_base(api_base)
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base)
if api_base is None or auth_header is None:
raise ValueError(
"ANTHROPIC_API_BASE/ANTHROPIC_BASE_URL or ANTHROPIC_API_KEY/ANTHROPIC_AUTH_TOKEN is not set. Please set the environment variable, to query Anthropic's `/models` endpoint."
)
headers: Final = {"anthropic-version": "2023-06-01"}
headers.update(auth_header)
response: Final = litellm.module_level_client.get(
url=f"{api_base}/v1/models",
headers=headers,
return self._list_models(api_key=api_key, api_base=api_base, litellm_params=None)
def discover_models(
self, litellm_params: Mapping[str, object] | None = None
) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override
"""Live discovery for a configured deployment: unlike ``get_models``, this threads the
full ``litellm_params`` into ``get_auth_header`` so a workload-identity-federation source
configured on the deployment (rather than the environment) is honored, gated the same way
every other Anthropic auth surface is via ``config_allows_workload_identity``."""
return self._list_models(
api_key=_litellm_params_str(litellm_params, "api_key"),
api_base=_litellm_params_str(litellm_params, "api_base"),
litellm_params=litellm_params,
)
try:
response.raise_for_status()
except httpx.HTTPStatusError:
raise Exception(
f"Failed to fetch models from Anthropic. Status code: {response.status_code}, Response: {response.text}"
def _list_models(
self,
*,
api_key: str | None,
api_base: str | None,
litellm_params: Mapping[str, object] | None,
) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override
resolved_api_base: Final = AnthropicModelInfo.get_api_base(api_base)
auth_header: Final = AnthropicModelInfo.get_auth_header(
api_key,
resolved_api_base,
litellm_params=litellm_params,
allow_workload_identity=config_allows_workload_identity(self),
)
if resolved_api_base is None or auth_header is None:
raise ValueError(
"ANTHROPIC_API_BASE/ANTHROPIC_BASE_URL or ANTHROPIC_API_KEY/ANTHROPIC_AUTH_TOKEN (or workload "
"identity federation via ANTHROPIC_FEDERATION_RULE_ID/ANTHROPIC_ORGANIZATION_ID/"
"ANTHROPIC_IDENTITY_TOKEN_FILE) is not set. Please set the environment variable, to query "
"Anthropic's `/models` endpoint."
)
models: Final[Sequence[Mapping[str, str]]] = response.json()["data"]
litellm_model_names: Final = ["anthropic/" + model["id"] for model in models]
return litellm_model_names
headers: Final = MappingProxyType({"anthropic-version": "2023-06-01", **auth_header})
# /v1/models is appended below, so a base the operator already wrote as .../v1 or
# .../v1/messages would otherwise be asked for /v1/v1/models.
model_ids: Final = _fetch_anthropic_model_ids(
anthropic_base_without_chat_suffix(resolved_api_base),
headers,
after_id=None,
pages_left=_MODEL_LIST_PAGE_CAP,
)
return ["anthropic/" + model_id for model_id in model_ids]
def get_token_counter(self) -> BaseTokenCounter | None:
"""
@ -1595,6 +1835,20 @@ def _replayed_server_tool_use(block: object) -> _ReplayedServerToolUse | None:
return None
def tool_call_is_rebuilt_as_server_tool_use(tool_call_id: object, provider_specific_fields: object) -> bool:
fields: Final = _validated_claude_code_mapping(provider_specific_fields)
if not isinstance(tool_call_id, str) or fields is None:
return False
return (
find_anthropic_server_tool_result(
tool_call_id,
_validated_claude_code_list(fields.get("web_search_results")),
_validated_claude_code_list(fields.get("tool_results")),
)
is not None
)
def _render_web_search_results(
query: str, results: tuple[_ReplayedWebSearchResult, ...] | _ReplayedWebSearchToolResultError
) -> str:

View file

@ -32,7 +32,7 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig):
self,
model: str,
messages: list[dict[str, JsonValue]],
api_key: str,
auth_header: Mapping[str, str],
api_base: str | None = None,
timeout: float | httpx.Timeout | None = None,
tools: list[dict[str, JsonValue]] | None = None,
@ -45,8 +45,8 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig):
Args:
model: The model identifier (e.g., "claude-3-5-sonnet-20241022")
messages: The messages to count tokens for
api_key: The Anthropic API key
api_base: Optional custom API base URL
auth_header: The resolved Anthropic auth header (``AnthropicModelInfo.get_auth_header``)
api_base: Optional deployment api_base the count-tokens path is appended to
timeout: Optional timeout for the request (defaults to litellm.request_timeout)
Returns:
@ -73,12 +73,12 @@ class AnthropicCountTokensHandler(AnthropicCountTokensConfig):
verbose_logger.debug("Transformed request: %s", request_body)
# Get endpoint URL
endpoint_url: Final = api_base or self.get_anthropic_count_tokens_endpoint()
endpoint_url: Final = self.get_anthropic_count_tokens_endpoint(api_base)
verbose_logger.debug("Making request to: %s", endpoint_url)
# Get required headers
headers: Final = self.get_required_headers(api_key)
headers: Final = self.get_count_tokens_headers(auth_header)
# Use LiteLLM's async httpx client
async_client: Final = get_async_httpx_client(llm_provider=litellm.LlmProviders.ANTHROPIC)

View file

@ -2,10 +2,10 @@
Anthropic Token Counter implementation using the CountTokens API.
"""
import os
from typing import Any, Final
from litellm._logging import verbose_logger
from litellm.exceptions import AuthenticationError
from litellm.llms.anthropic.count_tokens.handler import AnthropicCountTokensHandler
from litellm.llms.base_llm.base_utils import BaseTokenCounter
from litellm.types.utils import LlmProviders, TokenCountResponse
@ -46,28 +46,31 @@ class AnthropicTokenCounter(BaseTokenCounter):
Returns:
TokenCountResponse with token count, or None if counting fails
"""
from litellm.llms.anthropic.common_utils import AnthropicError
from litellm.llms.anthropic.common_utils import AnthropicError, AnthropicModelInfo
if not messages:
return None
deployment = deployment or {}
litellm_params: Final = deployment.get("litellm_params", {})
# Get Anthropic API key from deployment config or environment
api_key = litellm_params.get("api_key")
if not api_key:
api_key = os.getenv("ANTHROPIC_API_KEY")
if not api_key:
verbose_logger.warning("No Anthropic API key found for token counting")
return None
api_base: Final = litellm_params.get("api_base")
try:
auth_header: Final = await AnthropicModelInfo.aget_auth_header(
api_key=litellm_params.get("api_key"),
api_base=api_base,
litellm_params=litellm_params,
allow_workload_identity=True,
)
if auth_header is None:
verbose_logger.warning("No Anthropic credential found for token counting")
return None
result: Final = await anthropic_count_tokens_handler.handle_count_tokens_request(
model=model_to_use,
messages=messages,
api_key=api_key,
auth_header=auth_header,
api_base=api_base,
tools=tools,
system=system,
)
@ -80,8 +83,8 @@ class AnthropicTokenCounter(BaseTokenCounter):
tokenizer_type="anthropic_api",
original_response=result,
)
except AnthropicError as e:
verbose_logger.warning("Anthropic CountTokens API error: status=%s, message=%s", e.status_code, e.message)
except (AnthropicError, AuthenticationError) as e:
verbose_logger.warning("Anthropic CountTokens error: status=%s, message=%s", e.status_code, e.message)
return TokenCountResponse(
total_tokens=0,
request_model=request_model,

View file

@ -11,6 +11,8 @@ from typing import Final
from pydantic import JsonValue, TypeAdapter
from litellm.constants import ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
from litellm.llms.anthropic.common_utils import merge_anthropic_beta_headers
from litellm.llms.anthropic.wif import resolve_anthropic_base
_COUNT_REQUEST: Final = TypeAdapter(dict[str, JsonValue])
COUNT_TOKEN_OPTION_NAMES: Final = ("thinking", "tool_choice", "output_config")
@ -26,14 +28,21 @@ class AnthropicCountTokensConfig:
- Response: {"input_tokens": <number>}
"""
def get_anthropic_count_tokens_endpoint(self) -> str:
def get_anthropic_count_tokens_endpoint(self, api_base: str | None = None) -> str:
"""
Get the Anthropic CountTokens API endpoint.
Args:
api_base: The deployment's api_base, which names the chat surface (a host, or a
base already carrying ``/v1`` or ``/v1/messages``); the count-tokens path is
appended to it, so it is never the full count-tokens URL. Unset or empty falls
back to ``ANTHROPIC_API_BASE`` / ``ANTHROPIC_BASE_URL`` and then Anthropic's
host, the same resolution chat and the federated exchange use
Returns:
The endpoint URL for the CountTokens API
"""
return "https://api.anthropic.com/v1/messages/count_tokens"
return resolve_anthropic_base(api_base) + "/v1/messages/count_tokens"
def transform_request_to_count_tokens(
self,
@ -64,28 +73,19 @@ class AnthropicCountTokensConfig:
)
)
def get_required_headers(self, api_key: str) -> dict[str, str]:
"""
Get the required headers for the CountTokens API.
Args:
api_key: The Anthropic API key
Returns:
Dictionary of required headers
"""
from litellm.llms.anthropic.common_utils import (
optionally_handle_anthropic_oauth,
)
headers: dict[str, str] = {
def get_count_tokens_headers(self, auth_header: Mapping[str, str]) -> dict[str, str]:
"""The count-tokens headers around a resolved Anthropic auth header
(``AnthropicModelInfo.get_auth_header``): x-api-key for a static key, an Authorization
bearer for ``ANTHROPIC_AUTH_TOKEN`` and for sk-ant-oat tokens, whose mandatory oauth beta
merges with the token-counting beta instead of replacing it."""
return {
"Content-Type": "application/json",
"x-api-key": api_key,
"anthropic-version": "2023-06-01",
"anthropic-beta": ANTHROPIC_TOKEN_COUNTING_BETA_VERSION,
**auth_header,
"anthropic-beta": merge_anthropic_beta_headers(
auth_header.get("anthropic-beta"), ANTHROPIC_TOKEN_COUNTING_BETA_VERSION
),
}
headers, _ = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
return headers
def validate_request(
self,

View file

@ -1,7 +1,7 @@
import asyncio
import json
import time
from collections.abc import Coroutine
from collections.abc import Coroutine, Mapping
from typing import Final
import httpx
@ -43,6 +43,7 @@ class AnthropicFilesHandler:
api_key: str | None = None,
timeout: float | httpx.Timeout = 600.0,
max_retries: int | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> HttpxBinaryResponseContent:
"""
Async: Retrieve file content from Anthropic.
@ -56,6 +57,7 @@ class AnthropicFilesHandler:
api_key: Anthropic API key
timeout: Request timeout
max_retries: Max retry attempts (unused for now)
litellm_params: Deployment params, so a named credential's federation settings reach the mint
Returns:
HttpxBinaryResponseContent: Binary content wrapped in compatible response format
@ -73,7 +75,9 @@ class AnthropicFilesHandler:
# Get Anthropic API credentials
api_base = self.anthropic_model_info.get_api_base(api_base)
auth_header: Final = self.anthropic_model_info.get_auth_header(api_key, api_base)
auth_header: Final = await self.anthropic_model_info.aget_auth_header(
api_key, api_base, litellm_params=litellm_params, allow_workload_identity=True
)
if auth_header is None:
raise ValueError("Missing Anthropic API Key")
@ -116,6 +120,7 @@ class AnthropicFilesHandler:
api_key: str | None = None,
timeout: float | httpx.Timeout = 600.0,
max_retries: int | None = None,
litellm_params: Mapping[str, object] | None = None,
) -> HttpxBinaryResponseContent | Coroutine[object, object, HttpxBinaryResponseContent]:
"""
Retrieve file content from Anthropic.
@ -130,6 +135,7 @@ class AnthropicFilesHandler:
api_key: Anthropic API key
timeout: Request timeout
max_retries: Max retry attempts (unused for now)
litellm_params: Deployment params, so a named credential's federation settings reach the mint
Returns:
HttpxBinaryResponseContent or Coroutine: Binary content wrapped in compatible response format
@ -139,7 +145,9 @@ class AnthropicFilesHandler:
file_content_request=file_content_request,
api_base=api_base,
api_key=api_key,
timeout=timeout,
max_retries=max_retries,
litellm_params=litellm_params,
)
else:
return asyncio.run(
@ -149,6 +157,7 @@ class AnthropicFilesHandler:
api_key=api_key,
timeout=timeout,
max_retries=max_retries,
litellm_params=litellm_params,
)
)

View file

@ -14,6 +14,7 @@ Anthropic Files API endpoints:
import calendar
import time
from collections.abc import Mapping
from typing import Final, cast
import httpx
@ -35,7 +36,12 @@ from litellm.types.llms.openai import (
)
from litellm.types.utils import LlmProviders
from ..common_utils import AnthropicError, AnthropicModelInfo
from ..common_utils import (
AnthropicError,
AnthropicModelInfo,
merge_anthropic_beta_headers,
without_caller_credential_headers,
)
ANTHROPIC_FILES_API_BASE: Final = "https://api.anthropic.com"
ANTHROPIC_FILES_BETA_HEADER: Final = "files-api-2025-04-14"
@ -94,21 +100,55 @@ class AnthropicFilesConfig(BaseFilesConfig):
api_key: str | None = None,
api_base: str | None = None,
) -> dict:
if api_base is None and isinstance(litellm_params, dict):
api_base = litellm_params.get("api_base")
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base)
params_mapping, resolved_api_base = self._resolve_params(litellm_params, api_base)
auth_header: Final = AnthropicModelInfo.get_auth_header(
api_key, resolved_api_base, litellm_params=params_mapping, allow_workload_identity=True
)
return self._finalize_headers(headers, auth_header)
async def avalidate_environment(
self,
headers: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
model: str,
messages: list, # mutable-ok: mirrors the sync validate_environment contract this overrides
optional_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
litellm_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
api_key: str | None = None,
api_base: str | None = None,
) -> dict: # mutable-ok: mirrors the sync validate_environment contract this overrides
"""Async counterpart of validate_environment: the WIF tier can block on a token
exchange POST, so async callers await it off the event loop."""
params_mapping, resolved_api_base = self._resolve_params(litellm_params, api_base)
auth_header: Final = await AnthropicModelInfo.aget_auth_header(
api_key, resolved_api_base, litellm_params=params_mapping, allow_workload_identity=True
)
return self._finalize_headers(headers, auth_header)
@staticmethod
def _resolve_params(
litellm_params: dict, api_base: str | None
) -> tuple[dict | None, str | None]: # mutable-ok: mirrors the sync validate_environment contract this overrides
params_mapping: Final = litellm_params if isinstance(litellm_params, dict) else None
if api_base is None and params_mapping is not None:
api_base = params_mapping.get("api_base")
return params_mapping, api_base
@staticmethod
def _finalize_headers(headers: dict, auth_header: Mapping[str, str] | None) -> dict: # mutable-ok: out-param
if auth_header is None:
raise ValueError(
"Anthropic API key is required. Set ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN environment variable or pass api_key parameter."
)
headers.update(
{
**auth_header,
"anthropic-version": "2023-06-01",
"anthropic-beta": ANTHROPIC_FILES_BETA_HEADER,
}
merged_beta: Final = merge_anthropic_beta_headers(
merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")),
ANTHROPIC_FILES_BETA_HEADER,
)
return headers
return {
**without_caller_credential_headers(headers),
**auth_header,
"anthropic-version": "2023-06-01",
"anthropic-beta": merged_beta,
}
def get_supported_openai_params(self, model: str) -> list[OpenAICreateFileRequestOptionalParams]:
return ["purpose"]

View file

@ -1,5 +1,5 @@
from collections.abc import AsyncIterator, Mapping, Sequence
from typing import Any, Final
from typing import Any, ClassVar, Final
import httpx
@ -25,6 +25,7 @@ from litellm.types.router import GenericLiteLLMParams
from ...common_utils import (
AnthropicError,
AnthropicModelInfo,
merge_anthropic_beta_headers,
optionally_handle_anthropic_oauth,
requires_native_compaction_beta,
strip_advisor_blocks_from_messages,
@ -38,6 +39,17 @@ from .mid_conversation_system import (
DEFAULT_ANTHROPIC_API_VERSION: Final = "2023-06-01"
_CALLER_CREDENTIAL_HEADERS: Final = frozenset({"x-api-key", "authorization"})
def _carries_caller_credential(headers: Mapping[str, str]) -> bool:
"""Whether the caller sent their own Anthropic credential, in which case this passthrough
honors it and never mints. Matched case-insensitively: an SDK caller passing ``X-Api-Key``
through extra_headers would otherwise slip the check and end up sending their key beside a
minted federation Bearer."""
return any(name.lower() in _CALLER_CREDENTIAL_HEADERS for name in headers)
DROP_UNSUPPORTED_ADAPTIVE_EFFORT_WARNING: Final = (
"Dropping adaptive `thinking`/`output_config.effort` for model=%s: the model "
"does not support extended thinking, or max_tokens is too small to fit the "
@ -55,6 +67,8 @@ def _messages_carry_output_config(messages: Sequence[object]) -> bool:
class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
_workload_identity_eligible: ClassVar[bool] = True
@property
def custom_llm_provider(self) -> str | None:
return "anthropic"
@ -256,33 +270,109 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig):
# Check for Anthropic OAuth token in Authorization header
headers, api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
header_names: Final = frozenset(name.lower() for name in headers)
if "x-api-key" not in header_names and "authorization" not in header_names:
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key)
if auth_header is None:
raise AuthenticationError(
message=(
"Missing Anthropic API Key - A call is being made to anthropic but no key is set "
"either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` "
"or `ANTHROPIC_AUTH_TOKEN` in your environment vars"
if not _carries_caller_credential(headers):
self._apply_env_auth_header(
headers,
self._require_auth_header(
AnthropicModelInfo.get_auth_header(
api_key,
api_base=api_base,
litellm_params=litellm_params,
allow_workload_identity=self._allows_workload_identity,
),
llm_provider=self._resolved_provider,
model=model,
)
headers.update(auth_header)
),
)
return self._finalize_messages_headers(headers, optional_params, messages), api_base
async def avalidate_anthropic_messages_environment(
self,
headers: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
model: str,
messages: list[Any], # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
optional_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
litellm_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
api_key: str | None = None,
api_base: str | None = None,
) -> tuple[dict, str | None]: # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
if type(self).validate_anthropic_messages_environment is not (
AnthropicMessagesConfig.validate_anthropic_messages_environment
):
# a subclass sync override must keep winning on the async path
return self.validate_anthropic_messages_environment(
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
)
oauth_headers, oauth_api_key = optionally_handle_anthropic_oauth(headers=headers, api_key=api_key)
if not _carries_caller_credential(oauth_headers):
self._apply_env_auth_header(
oauth_headers,
self._require_auth_header(
await AnthropicModelInfo.aget_auth_header(
oauth_api_key,
api_base=api_base,
litellm_params=litellm_params,
allow_workload_identity=self._allows_workload_identity,
),
model=model,
),
)
return self._finalize_messages_headers(oauth_headers, optional_params, messages), api_base
def _require_auth_header(self, auth_header: Mapping[str, str] | None, model: str) -> Mapping[str, str]:
if auth_header is None:
raise AuthenticationError(
message=(
"Missing Anthropic API Key - A call is being made to anthropic but no key is set "
"either in the environment variables or via params. Please set `ANTHROPIC_API_KEY` "
"or `ANTHROPIC_AUTH_TOKEN` in your environment vars"
),
llm_provider=self._resolved_provider,
model=model,
)
return auth_header
@staticmethod
def _apply_env_auth_header(headers: dict, auth_header: Mapping[str, str] | None) -> None: # mutable-ok: out-param
if auth_header is None:
return
merged_beta: Final = merge_anthropic_beta_headers(
headers.get("anthropic-beta"), auth_header.get("anthropic-beta")
)
headers.update(auth_header)
if merged_beta:
headers["anthropic-beta"] = merged_beta
@property
def _allows_workload_identity(self) -> bool:
"""Subclasses reuse this validate step for their own /v1/messages-compatible providers, so
eligibility is declared per class and never inherited."""
from litellm.llms.anthropic.common_utils import config_allows_workload_identity
return config_allows_workload_identity(self)
def _finalize_messages_headers(
self,
headers: dict, # mutable-ok: out-param
optional_params: dict, # mutable-ok: out-param
messages: list[Any], # mutable-ok: mirrors the validate_anthropic_messages_environment contract
) -> dict: # mutable-ok: out-param
if "anthropic-version" not in headers:
headers["anthropic-version"] = DEFAULT_ANTHROPIC_API_VERSION
if "content-type" not in headers:
headers["content-type"] = "application/json"
headers = self._update_headers_with_anthropic_beta(
return self._update_headers_with_anthropic_beta(
headers=headers,
optional_params=optional_params,
messages=messages,
)
return headers, api_base
@staticmethod
def _translate_reasoning_effort_to_anthropic(
model: str, optional_params: dict, max_tokens: int | None, custom_llm_provider: str

View file

@ -524,17 +524,19 @@ async def count_prompt_tokens(
body: Mapping[str, JsonValue],
api_base: str | None = None,
) -> int | None:
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key=api_key, api_base=api_base)
if auth_header is None:
return None
try:
native: Final = _CountBody.model_validate(body)
count_url: Final = _messages_url(model, api_key, api_base) + "/count_tokens"
result: Final = _CountResult.model_validate(
await _counter.handle_count_tokens_request(
model=model,
messages=_count_objects(native.messages),
tools=_count_objects(native.tools) if native.tools is not None else None,
system=_JSON_OBJECT.validate_python(MappingProxyType({"system": native.system}))["system"],
api_key=api_key,
api_base=count_url,
auth_header=auth_header,
api_base=api_base,
optional_params=_JSON_OBJECT.validate_python(
MappingProxyType({key: body[key] for key in COUNT_TOKEN_OPTION_NAMES if key in body})
),

View file

@ -2,6 +2,7 @@
Anthropic Skills API configuration and transformations
"""
from types import MappingProxyType
from typing import Final
import httpx
@ -35,40 +36,35 @@ class AnthropicSkillsConfig(BaseSkillsAPIConfig):
def validate_environment(self, headers: dict, litellm_params: GenericLiteLLMParams | None) -> dict:
"""Add Anthropic-specific headers"""
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.constants import ANTHROPIC_SKILLS_API_BETA_VERSION
from litellm.llms.anthropic.common_utils import (
AnthropicModelInfo,
merge_anthropic_beta_headers,
without_caller_credential_headers,
)
# Get API key from litellm_params if available
api_key = None
api_base = None
if litellm_params is not None:
api_key = litellm_params.api_key
api_base = litellm_params.api_base
auth_header: Final = AnthropicModelInfo.get_auth_header(api_key, api_base)
auth_header: Final = AnthropicModelInfo.get_auth_header(
api_key=litellm_params.api_key if litellm_params is not None else None,
api_base=litellm_params.api_base if litellm_params is not None else None,
litellm_params=MappingProxyType(dict(litellm_params)) if litellm_params is not None else None,
allow_workload_identity=True,
)
if auth_header is None:
raise ValueError("ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN is required for Skills API")
headers.update(auth_header)
headers["anthropic-version"] = "2023-06-01"
# Add beta header for skills API
from litellm.constants import ANTHROPIC_SKILLS_API_BETA_VERSION
if "anthropic-beta" not in headers:
headers["anthropic-beta"] = ANTHROPIC_SKILLS_API_BETA_VERSION
elif isinstance(headers["anthropic-beta"], list):
if ANTHROPIC_SKILLS_API_BETA_VERSION not in headers["anthropic-beta"]:
headers["anthropic-beta"].append(ANTHROPIC_SKILLS_API_BETA_VERSION)
elif isinstance(headers["anthropic-beta"], str):
if ANTHROPIC_SKILLS_API_BETA_VERSION not in headers["anthropic-beta"]:
headers["anthropic-beta"] = [
headers["anthropic-beta"],
ANTHROPIC_SKILLS_API_BETA_VERSION,
]
headers["content-type"] = "application/json"
return headers
merged_beta: Final = merge_anthropic_beta_headers(
merge_anthropic_beta_headers(headers.get("anthropic-beta"), auth_header.get("anthropic-beta")),
ANTHROPIC_SKILLS_API_BETA_VERSION,
)
# The deployment's own credential is applied here, so a caller-supplied one must not ride
# along upstream beside a minted federation Bearer.
return {
**without_caller_credential_headers(headers),
**auth_header,
"anthropic-version": "2023-06-01",
"anthropic-beta": merged_beta,
"content-type": "application/json",
}
def get_complete_url(
self,

View file

@ -0,0 +1,604 @@
"""Anthropic workload identity federation: exchanges an external OIDC identity
token for a short-lived ``sk-ant-oat01`` token via the shared RFC 7523 engine."""
import os
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from functools import lru_cache
from itertools import chain
from types import MappingProxyType
from typing import Final, NoReturn, TypeVar
from urllib.parse import urlsplit, urlunsplit
from pydantic import BaseModel, ConfigDict, ValidationError
from typing_extensions import assert_never
import litellm
from litellm._logging import verbose_logger
from litellm.llms.base_llm.auth.client_credentials import keycloak_assertion_source
from litellm.llms.base_llm.auth.identity_source import (
AnthropicIdentitySourceKind,
InternalIssuerSource,
KeycloakSource,
identity_source_ref,
)
from litellm.llms.base_llm.auth.internal_issuer import (
internal_issuer_assertion_source,
internal_issuer_jwks_document,
)
from litellm.llms.base_llm.auth.token_exchange import (
JwtBearerTokenExchangeEngine,
default_token_exchange_engine,
)
from litellm.llms.base_llm.auth.types import (
AssertionSourceError,
ExchangeError,
ExchangeResult,
InsecureTokenUrl,
MalformedTokenResponse,
MintedToken,
TokenEndpointError,
TokenExchangeSpec,
TokenTransportError,
)
from litellm.types.llms.anthropic import ANTHROPIC_TOKEN_EXCHANGE_PATH
_JWT_BEARER_GRANT_TYPE: Final = "urn:ietf:params:oauth:grant-type:jwt-bearer"
_DEFAULT_API_BASE: Final = "https://api.anthropic.com"
_INLINE_ENV_VAR: Final = "ANTHROPIC_IDENTITY_TOKEN"
_DISABLE_WIF_PARAM: Final = "anthropic_disable_workload_identity_federation"
_ACCEPTED_REF_PREFIX: Final = "oidc/"
_SHADOWED_DEPLOYMENT_WARNING_CAP: Final = 512
_CHAT_BASE_SUFFIXES: Final = ("/v1/messages", "/v1")
# Hosts a federated exchange may talk to. api_base decides where the workload's assertion is sent
# AND where the minted org-scoped token is presented, so anyone able to write api_base on a
# federated deployment could otherwise redirect both. Gating each write path does not terminate:
# a deployment, a referenced credential and a future endpoint all reach the same value. This is the
# one place a federated exchange is built, so the trust decision is enforced here instead, and the
# allowlist is server-owned -- read from the environment, never from a model or credential API.
_TRUSTED_EXCHANGE_HOSTS_ENV: Final = "LITELLM_ANTHROPIC_WIF_ALLOWED_HOSTS"
_SCHEME_DEFAULT_PORTS: Final[Mapping[str, int]] = MappingProxyType({"http": 80, "https": 443})
_DEFAULT_TRUSTED_EXCHANGE_HOST: Final = "api.anthropic.com"
_REJECTED_REF_PREFIX: Final = "oidc/env_path/"
_IDENTITY_SOURCE_PARAM: Final = "anthropic_identity_source"
_IDENTITY_SOURCE_ENV: Final = "ANTHROPIC_IDENTITY_SOURCE"
_IDENTITY_TOKEN_FILE_PARAM: Final = "anthropic_identity_token_file"
_IDENTITY_TOKEN_PARAM: Final = "anthropic_identity_token"
# litellm_params key -> InternalIssuerSource/KeycloakSource field name. Every key here must
# also be listed in ANTHROPIC_WIF_KWARGS_KEYS (types/workload_identity.py), which is what makes it
# request-banned and cleared on a client-redirected api_base -- see types/utils.py's
# anthropic_wif_litellm_params, derived from that same set.
_INTERNAL_ISSUER_FIELD_MAP: Final[Mapping[str, str]] = MappingProxyType(
{
"anthropic_issuer_url": "issuer_url",
"anthropic_issuer_subject": "subject",
"anthropic_issuer_audience": "audience",
"anthropic_issuer_ttl_seconds": "ttl_seconds",
"anthropic_issuer_signing_key_ref": "signing_key_ref",
}
)
_KEYCLOAK_FIELD_MAP: Final[Mapping[str, str]] = MappingProxyType(
{
"anthropic_keycloak_token_url": "token_url",
"anthropic_keycloak_client_id": "client_id",
"anthropic_keycloak_auth_method": "auth_method",
"anthropic_keycloak_client_secret_ref": "client_secret_ref",
"anthropic_keycloak_scope": "scope",
}
)
_DENIAL_HINT: Final = (
"Anthropic answers every denied exchange with the same 401; the reason (for example"
" workspace_id_required or jti_reused) is only shown in the Claude Console under"
" Settings > Workload identity, in the rule's authentication history. jti_reused means this"
" identity token was already exchanged once: Anthropic accepts each assertion a single time, so a"
" token file or env var has to rotate before the minted token expires (the rule's"
" token_lifetime_seconds), or switch to the internal issuer or Keycloak source, which mint a"
" fresh assertion per exchange"
)
_WORKSPACE_HINT: Final = (
"If the federation rule is enabled in more than one workspace, set anthropic_federation_workspace_id"
" (or ANTHROPIC_FEDERATION_WORKSPACE_ID) to the wrkspc_ id of the workspace to mint tokens for, or to 'default'."
" Federation does not read ANTHROPIC_WORKSPACE_ID, which the Bedrock Claude platform provider already uses"
)
_SERVICE_ACCOUNT_HINT: Final = (
"Anthropic's reference lists service_account_id as required: set anthropic_service_account_id"
" (or ANTHROPIC_SERVICE_ACCOUNT_ID) to the svac_ id the federation rule targets"
)
_MISSING_IDS_HINT: Final = (
"Copy them from the federation rule's detail page under Settings > Workload identity in the"
" Claude Console, or set ANTHROPIC_FEDERATION_RULE_ID and ANTHROPIC_ORGANIZATION_ID"
)
_ALLOWLIST_HINT: Final = (
"Identity token files must sit under an allowed credential directory"
" (/var/run/secrets or /run/secrets by default);"
" set LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS to extend the allowlist"
)
_EMPTY_PARAMS: Final[Mapping[str, object]] = MappingProxyType({})
_IdentitySourceVariant = TypeVar("_IdentitySourceVariant", bound="InternalIssuerSource | KeycloakSource")
class AnthropicWifParams(BaseModel):
model_config = ConfigDict(frozen=True)
federation_rule_id: str
organization_id: str
service_account_id: str | None = None
workspace_id: str | None = None
assertion_ref: str
assertion_source: Callable[[], str | None] | None = None
def resolve_anthropic_wif_params(litellm_params: Mapping[str, object] | None) -> AnthropicWifParams | None:
if litellm_params is not None and litellm_params.get(_DISABLE_WIF_PARAM) is True:
return None
federation_rule_id: Final = _config_value(
litellm_params, "anthropic_federation_rule_id", "ANTHROPIC_FEDERATION_RULE_ID"
)
organization_id: Final = _config_value(litellm_params, "anthropic_organization_id", "ANTHROPIC_ORGANIZATION_ID")
if federation_rule_id is None or organization_id is None:
_raise_if_identity_source_configured(litellm_params, federation_rule_id, organization_id)
return None
identity_source: Final = _resolve_identity_source(litellm_params)
if identity_source is None:
return None
assertion_ref, assertion_source = identity_source
return AnthropicWifParams(
federation_rule_id=federation_rule_id,
organization_id=organization_id,
service_account_id=_config_value(
litellm_params, "anthropic_service_account_id", "ANTHROPIC_SERVICE_ACCOUNT_ID"
),
workspace_id=_config_value(
litellm_params, "anthropic_federation_workspace_id", "ANTHROPIC_FEDERATION_WORKSPACE_ID"
),
assertion_ref=assertion_ref,
assertion_source=assertion_source,
)
def _resolve_identity_source(
litellm_params: Mapping[str, object] | None,
) -> tuple[str, Callable[[], str] | None] | None:
"""Dispatches on ``anthropic_identity_source``. Absent (the default) keeps today's
token_file/env resolution byte-identical, with no ``assertion_source`` closure -- the engine
falls back to its own reader exactly as it does today. A recognized kind builds the matching
frozen config, hashes it into the ``oidc/<kind>/<hash>`` cache-key ref (``identity_source_ref``),
and closes the source's fetch/mint function over it. An unset-but-invalid config (unknown
kind, a missing required field, or a field from the other variant) fails closed here rather
than silently falling back to token_file. A deployment whose params carry a legacy token or
token_file ref stays on legacy resolution even when ``ANTHROPIC_IDENTITY_SOURCE`` names a
fleet-wide kind: the env kind only governs deployments that set no identity params of their own."""
source_kind: Final = _resolve_source_kind(litellm_params)
if source_kind is None:
legacy_ref: Final = _resolve_assertion_ref(litellm_params)
return (legacy_ref, None) if legacy_ref is not None else None
params: Final[Mapping[str, object]] = MappingProxyType(
{key: value for key, value in (litellm_params or _EMPTY_PARAMS).items() if _is_set(value)}
)
match source_kind:
case AnthropicIdentitySourceKind.internal_issuer.value:
_reject_foreign_variant_fields(params, foreign_field_map=_KEYCLOAK_FIELD_MAP, chosen_kind=source_kind)
issuer_config: Final = _build_variant(InternalIssuerSource, params, _INTERNAL_ISSUER_FIELD_MAP)
return identity_source_ref(issuer_config), internal_issuer_assertion_source(issuer_config)
case AnthropicIdentitySourceKind.keycloak.value:
_reject_foreign_variant_fields(
params, foreign_field_map=_INTERNAL_ISSUER_FIELD_MAP, chosen_kind=source_kind
)
keycloak_config: Final = _build_variant(KeycloakSource, params, _KEYCLOAK_FIELD_MAP)
return identity_source_ref(keycloak_config), keycloak_assertion_source(keycloak_config)
case _:
_raise_unknown_source_kind(source_kind)
def _raise_unknown_source_kind(source_kind: str) -> NoReturn:
raise litellm.AuthenticationError(
message=(
f"{_IDENTITY_SOURCE_PARAM} must be one of "
f"{', '.join(kind.value for kind in AnthropicIdentitySourceKind)}; got {source_kind!r}"
),
llm_provider="anthropic",
model="",
)
def _raise_if_identity_source_configured(
litellm_params: Mapping[str, object] | None, federation_rule_id: str | None, organization_id: str | None
) -> None:
"""A configured identity source is an explicit request to federate, so a missing rule or
organization id fails closed with the ids named, rather than silently skipping federation
and surfacing later as a missing API key."""
source_kind: Final = _resolve_source_kind(litellm_params)
if source_kind is None:
return
if source_kind not in {kind.value for kind in AnthropicIdentitySourceKind}:
_raise_unknown_source_kind(source_kind)
missing: Final = tuple(
param
for param, value in (
("anthropic_federation_rule_id", federation_rule_id),
("anthropic_organization_id", organization_id),
)
if value is None
)
raise litellm.AuthenticationError(
message=(
f"{_IDENTITY_SOURCE_PARAM} is {source_kind!r}, but {' and '.join(missing)} "
f"{'is' if len(missing) == 1 else 'are'} not set. {_MISSING_IDS_HINT}"
),
llm_provider="anthropic",
model="",
)
def _resolve_source_kind(litellm_params: Mapping[str, object] | None) -> str | None:
param_kind: Final = _param_str(litellm_params, _IDENTITY_SOURCE_PARAM)
if param_kind is not None:
return param_kind
has_param_legacy_ref: Final = any(
_param_str(litellm_params, key) is not None for key in (_IDENTITY_TOKEN_FILE_PARAM, _IDENTITY_TOKEN_PARAM)
)
return None if has_param_legacy_ref else _env_str(_IDENTITY_SOURCE_ENV)
def _reject_foreign_variant_fields(
litellm_params: Mapping[str, object], foreign_field_map: Mapping[str, str], chosen_kind: str
) -> None:
foreign_keys_present: Final = tuple(param for param in foreign_field_map if param in litellm_params)
if foreign_keys_present:
raise litellm.AuthenticationError(
message=(
f"{_IDENTITY_SOURCE_PARAM} is {chosen_kind!r}, but {', '.join(sorted(foreign_keys_present))} "
"belongs to a different identity source and cannot be set alongside it"
),
llm_provider="anthropic",
model="",
)
def _build_variant(
model: type[_IdentitySourceVariant],
litellm_params: Mapping[str, object],
field_map: Mapping[str, str],
) -> _IdentitySourceVariant:
fields: Final = MappingProxyType(
{field_map[key]: value for key, value in litellm_params.items() if key in field_map and _is_set(value)}
)
try:
return model.model_validate(fields)
except ValidationError as e:
# hide_input_in_errors=True on both variant models keeps a secret pasted into the
# wrong field (e.g. a client_secret typed as signing_key_ref) out of str(e).
raise litellm.AuthenticationError(
message=f"Invalid {_IDENTITY_SOURCE_PARAM} configuration: {e}",
llm_provider="anthropic",
model="",
) from e
@dataclass(frozen=True, slots=True)
class ExportedJwks:
document: str
@dataclass(frozen=True, slots=True)
class NotAnInternalIssuerCredential:
required_param: str
required_value: str
@dataclass(frozen=True, slots=True)
class UnbuildableIdentitySource:
message: str
AnthropicJwksExport = ExportedJwks | NotAnInternalIssuerCredential | UnbuildableIdentitySource
def anthropic_internal_issuer_jwks(credential_values: Mapping[str, object]) -> AnthropicJwksExport:
"""Derive the public JWKS a stored anthropic credential publishes to its federation issuer.
The private signing key stays in this process; only the derived public document comes back."""
if credential_values.get(_IDENTITY_SOURCE_PARAM) != AnthropicIdentitySourceKind.internal_issuer.value:
return NotAnInternalIssuerCredential(
required_param=_IDENTITY_SOURCE_PARAM,
required_value=AnthropicIdentitySourceKind.internal_issuer.value,
)
try:
issuer_source: Final = _build_variant(InternalIssuerSource, credential_values, _INTERNAL_ISSUER_FIELD_MAP)
return ExportedJwks(internal_issuer_jwks_document(issuer_source))
except (litellm.AuthenticationError, ValueError) as e:
return UnbuildableIdentitySource(str(e))
def build_anthropic_wif_spec(params: AnthropicWifParams, api_base: str) -> TokenExchangeSpec:
return TokenExchangeSpec(
token_url=api_base.rstrip("/") + ANTHROPIC_TOKEN_EXCHANGE_PATH,
assertion_ref=params.assertion_ref,
assertion_field="assertion",
static_body=MappingProxyType(
{
name: value
for name, value in (
("grant_type", _JWT_BEARER_GRANT_TYPE),
("federation_rule_id", params.federation_rule_id),
("organization_id", params.organization_id),
("service_account_id", params.service_account_id),
("workspace_id", params.workspace_id),
)
if value is not None
}
),
body_encoding="json",
request_headers=MappingProxyType({}),
assertion_source=params.assertion_source,
cache_key_identity=(
params.federation_rule_id,
params.organization_id,
params.service_account_id or "",
params.workspace_id or "",
),
)
def get_anthropic_wif_token(
litellm_params: Mapping[str, object] | None,
api_base: str | None,
model: str,
engine: JwtBearerTokenExchangeEngine = default_token_exchange_engine,
) -> str | None:
params: Final = resolve_anthropic_wif_params(litellm_params)
if params is None:
return None
exchange_base: Final = resolve_anthropic_base(api_base)
_raise_if_exchange_host_untrusted(exchange_base, model)
result: Final = engine.get_token(build_anthropic_wif_spec(params, exchange_base))
return _token_from_result(result, model, params)
async def aget_anthropic_wif_token(
litellm_params: Mapping[str, object] | None,
api_base: str | None,
model: str,
engine: JwtBearerTokenExchangeEngine = default_token_exchange_engine,
) -> str | None:
params: Final = resolve_anthropic_wif_params(litellm_params)
if params is None:
return None
exchange_base: Final = resolve_anthropic_base(api_base)
_raise_if_exchange_host_untrusted(exchange_base, model)
result: Final = await engine.aget_token(build_anthropic_wif_spec(params, exchange_base))
return _token_from_result(result, model, params)
def _token_from_result(result: ExchangeResult, model: str, params: AnthropicWifParams) -> str:
match result:
case MintedToken():
return result.access_token.get_secret_value()
case _:
_raise_anthropic_wif_error(
result,
model=model,
workspace_id_set=params.workspace_id is not None,
service_account_id_set=params.service_account_id is not None,
)
def resolve_anthropic_base(api_base: str | None) -> str:
"""The base every Anthropic tier derives its URLs from: the deployment api_base when set,
else ``ANTHROPIC_API_BASE`` / ``ANTHROPIC_BASE_URL``, else Anthropic's host, with trailing
slashes and chat-appended ``/v1/messages`` suffixes stripped, so the token URL, the cache key
and the count-tokens URL all agree for the same deployment."""
return anthropic_base_without_chat_suffix(api_base or _resolve_default_api_base())
def _allowlisted_authority(entry: str) -> tuple[str, int | None] | None:
"""One allowlist entry as ``(host, port)``. The port stays ``None`` unless the entry spells one
out, so ``gateway.internal`` trusts that host on every port while ``gateway.internal:8443``
trusts only 8443."""
parts: Final = urlsplit(entry if "://" in entry else f"//{entry}")
try:
port: Final = parts.port
except ValueError:
return None
return (parts.hostname, port) if parts.hostname else None
def _exchange_authority(exchange_base: str) -> tuple[str, int | None]:
"""The host and port an exchange would actually reach, filling in the scheme's default port so
an operator who wrote ``api.anthropic.com:443`` still matches ``https://api.anthropic.com``."""
parts: Final = urlsplit(exchange_base)
try:
port: Final = parts.port
except ValueError:
return "", None
return (parts.hostname or "").lower(), port if port is not None else _SCHEME_DEFAULT_PORTS.get(parts.scheme)
def _trusted_exchange_authorities() -> frozenset[tuple[str, int | None]]:
"""Authorities a federated exchange may reach: Anthropic's own, plus whatever the operator put in
the environment. Comma separated, case folded, each entry a URL, a bare host, or ``host:port``."""
configured: Final = os.getenv(_TRUSTED_EXCHANGE_HOSTS_ENV) or ""
entries: Final = (_allowlisted_authority(entry.strip()) for entry in configured.split(",") if entry.strip())
return frozenset(chain(((_DEFAULT_TRUSTED_EXCHANGE_HOST, None),), (entry for entry in entries if entry)))
def _raise_if_exchange_host_untrusted(exchange_base: str, model: str) -> None:
"""The federated exchange refuses any authority the operator has not vouched for, whatever wrote
the deployment's api_base. Exact host match, never a substring: ``api.anthropic.com.evil.test``
contains the real host and must not pass. An entry naming a port trusts that port alone, so a
second process on another port of an allowed host is refused."""
host, port = _exchange_authority(exchange_base)
if host and any(
host == allowed_host and allowed_port in (None, port)
for allowed_host, allowed_port in _trusted_exchange_authorities()
):
return
refused: Final = f"{host}:{port}" if host and port is not None else host
raise litellm.AuthenticationError(
message=(
f"Anthropic workload identity federation refused to use host {refused or exchange_base!r}. "
f"A federated exchange sends the workload's identity token to this host and presents the "
f"minted token to it, so only {_DEFAULT_TRUSTED_EXCHANGE_HOST} is trusted by default. To "
f"use a private Anthropic-compatible gateway, add its host, or host:port to pin the port, "
f"to the {_TRUSTED_EXCHANGE_HOSTS_ENV} environment variable (comma separated); that is a "
f"decision to trust it with org-scoped credentials, so it is deliberately server-owned "
f"and cannot be set through the model or credential APIs"
),
llm_provider="anthropic",
model=model,
)
@lru_cache(maxsize=_SHADOWED_DEPLOYMENT_WARNING_CAP)
def _warn_static_credential_shadows_federation(model: str, configured_rule_id: str | None) -> None:
"""Memoized so a shadowed deployment says this once rather than once per request.
The environment fallback is resolved in here rather than by the caller so it too costs one
secret-manager read per deployment: every static-key Anthropic call reaches this, and a
per-request read of a rule id almost nobody sets is an ERROR log with a traceback per call on
the deployments that shadow nothing.
"""
if configured_rule_id is None and _env_str("ANTHROPIC_FEDERATION_RULE_ID") is None:
return
verbose_logger.warning(
"Anthropic deployment %s is configured for workload identity federation, but a static "
"ANTHROPIC_API_KEY or ANTHROPIC_AUTH_TOKEN is set and takes precedence, so every call bills "
"that credential and no federated token is minted. Unset it to federate.",
model or "(unnamed)",
)
def warn_if_static_credential_shadows_federation(litellm_params: Mapping[str, object] | None, model: str) -> None:
"""A process-wide static credential outranks federation everywhere in the provider, which is the
Anthropic SDK's own precedence. An operator who configured federation and left a key behind would
otherwise get no signal at all that none of their calls are federated."""
if litellm_params is not None and litellm_params.get(_DISABLE_WIF_PARAM) is True:
return
_warn_static_credential_shadows_federation(model, _param_str(litellm_params, "anthropic_federation_rule_id"))
def _resolve_default_api_base() -> str:
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
return AnthropicModelInfo.get_api_base(None) or _DEFAULT_API_BASE
def anthropic_base_without_chat_suffix(base: str) -> str:
"""A deployment base with its chat-surface suffix removed, so the token URL and model
discovery both derive from the same value whatever form the operator configured."""
parts: Final = urlsplit(base)
if not parts.scheme or not parts.netloc:
return base.rstrip("/")
return urlunsplit((parts.scheme, parts.netloc, _strip_path_suffixes(parts.path), "", ""))
def _strip_path_suffixes(path: str) -> str:
"""Drop the chat-surface suffixes a deployment base may carry, so every tier derives the same
token URL. Each pass removes at most one suffix, so the loop is bounded by the segment count."""
trimmed = path.rstrip("/") # rebind-ok: fixed-point strip, one suffix per pass
while True:
shortened = next(
(trimmed.removesuffix(suffix) for suffix in _CHAT_BASE_SUFFIXES if trimmed.endswith(suffix)),
trimmed,
)
if shortened == trimmed:
return trimmed
# Re-strip: a doubled suffix leaves a trailing slash that would stop the next match.
trimmed = shortened.rstrip("/")
def _config_value(litellm_params: Mapping[str, object] | None, param_key: str, env_name: str) -> str | None:
return _param_str(litellm_params, param_key) or _env_str(env_name)
def _is_set(value: object) -> bool:
return value is not None and value != ""
def _param_str(litellm_params: Mapping[str, object] | None, key: str) -> str | None:
if litellm_params is None:
return None
value: Final = litellm_params.get(key)
return value if isinstance(value, str) and value else None
def _env_str(name: str) -> str | None:
from litellm.secret_managers.main import get_secret_str
value: Final = get_secret_str(name)
return value if isinstance(value, str) and value else None
def _resolve_assertion_ref(litellm_params: Mapping[str, object] | None) -> str | None:
file_param: Final = _param_str(litellm_params, _IDENTITY_TOKEN_FILE_PARAM)
if file_param is not None:
return f"oidc/file/{file_param}"
inline_param: Final = _param_str(litellm_params, _IDENTITY_TOKEN_PARAM)
if inline_param is not None:
return _validated_inline_ref(inline_param)
file_env: Final = _env_str("ANTHROPIC_IDENTITY_TOKEN_FILE")
if file_env is not None:
return f"oidc/file/{file_env}"
if _env_str(_INLINE_ENV_VAR) is not None:
return f"oidc/env/{_INLINE_ENV_VAR}"
return None
def _validated_inline_ref(value: str) -> str:
if value.startswith(_ACCEPTED_REF_PREFIX) and not value.startswith(_REJECTED_REF_PREFIX):
return value
raise litellm.AuthenticationError(
message=(
"anthropic_identity_token must be an oidc/ secret reference such as oidc/env/VAR_NAME,"
" oidc/file//absolute/path, oidc/github/<audience>, or oidc/google/<audience>."
" Raw identity tokens and oidc/env_path/ references are not accepted;"
" to pass a token directly, export it and reference it as oidc/env/VAR_NAME"
),
llm_provider="anthropic",
model="",
)
def _raise_anthropic_wif_error(
error: ExchangeError, model: str, workspace_id_set: bool, service_account_id_set: bool
) -> NoReturn:
detail: Final = _error_detail(
error, workspace_id_set=workspace_id_set, service_account_id_set=service_account_id_set
)
raise litellm.AuthenticationError(
message=f"Anthropic workload identity federation failed. {detail}",
llm_provider="anthropic",
model=model,
)
def _denial_hints(workspace_id_set: bool, service_account_id_set: bool) -> str:
hints: Final = (
_DENIAL_HINT,
"" if workspace_id_set else _WORKSPACE_HINT,
"" if service_account_id_set else _SERVICE_ACCOUNT_HINT,
)
return " " + ". ".join(hint for hint in hints if hint)
def _error_detail(error: ExchangeError, workspace_id_set: bool, service_account_id_set: bool) -> str:
match error:
case AssertionSourceError() if error.kind == "disallowed_path":
return f"Could not read the OIDC identity token from {error.source_ref}. {_ALLOWLIST_HINT}"
case AssertionSourceError():
base: Final = f"Could not obtain the OIDC identity token ({error.kind}) from {error.source_ref}"
return f"{base}. {error.detail}" if error.detail else base
case InsecureTokenUrl():
return f"The token endpoint must use https; refusing to send the identity token to host {error.host!r}"
case TokenEndpointError() if error.status_code == 401:
hints: Final = _denial_hints(workspace_id_set, service_account_id_set)
return f"The token endpoint returned HTTP 401: {error.redacted_body}{hints}"
case TokenEndpointError():
return f"The token endpoint returned HTTP {error.status_code}: {error.redacted_body}"
case TokenTransportError():
return f"Could not reach the token endpoint: {error.detail}"
case MalformedTokenResponse():
return f"The token endpoint returned an unusable response: {error.detail}"
case _:
assert_never(error)

View file

@ -1,3 +1,4 @@
from collections.abc import Mapping
from typing import Final
from urllib.parse import urlsplit, urlunsplit
@ -218,6 +219,7 @@ class AzureAIEmbedding(OpenAIChatCompletion):
aembedding=None,
max_retries: int | None = None,
shared_session=None,
litellm_params: Mapping[str, object] | None = None,
) -> EmbeddingResponse:
"""
- Separate image url from text

View file

@ -41,6 +41,29 @@ class BaseAnthropicMessagesConfig(ABC):
"""
return headers, api_base
async def avalidate_anthropic_messages_environment(
self,
headers: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
model: str,
messages: list[Any], # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
optional_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
litellm_params: dict, # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
api_key: str | None = None,
api_base: str | None = None,
) -> tuple[dict, str | None]: # mutable-ok: mirrors the sync validate_anthropic_messages_environment contract
"""Async counterpart used by the async handler. The default delegates to the
sync implementation; providers whose sync path can block the event loop
(e.g. a WIF token exchange) override this."""
return self.validate_anthropic_messages_environment(
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
api_base=api_base,
)
@abstractmethod
def get_complete_url(
self,

View file

@ -0,0 +1,99 @@
from litellm.llms.base_llm.auth.client_credentials import (
SecretReader,
fetch_keycloak_assertion,
keycloak_assertion_source,
)
from litellm.llms.base_llm.auth.identity_source import (
AnthropicIdentitySourceConfig,
AnthropicIdentitySourceKind,
InternalIssuerSource,
KeycloakSource,
identity_source_config_adapter,
identity_source_ref,
)
from litellm.llms.base_llm.auth.internal_issuer import (
SigningKeyReader,
internal_issuer_assertion_source,
internal_issuer_jwks_document,
mint_internal_issuer_assertion,
)
from litellm.llms.base_llm.auth.jwt_signing import (
ALG,
build_jwk,
build_jwks,
jwks_document_json,
load_es256_private_key,
rfc7638_thumbprint,
sign_es256_jwt,
)
from litellm.llms.base_llm.auth.token_exchange import (
ADVISORY_REFRESH_BACKOFF_SECONDS,
ADVISORY_REFRESH_SECONDS,
MANDATORY_REFRESH_SECONDS,
MAX_ASSERTION_BYTES,
MAX_RESPONSE_BYTES,
JwtBearerTokenExchangeEngine,
default_token_exchange_engine,
redact_oauth_error_body,
validate_token_endpoint_url,
)
from litellm.llms.base_llm.auth.types import (
AssertionReader,
AssertionSource,
AssertionSourceError,
BodyEncoding,
ExchangeError,
ExchangeResult,
InsecureTokenUrl,
MalformedTokenResponse,
MintedToken,
SyncTokenPoster,
TokenEndpointError,
TokenExchangeSpec,
TokenTransportError,
)
__all__ = (
"ADVISORY_REFRESH_BACKOFF_SECONDS",
"ADVISORY_REFRESH_SECONDS",
"ALG",
"MANDATORY_REFRESH_SECONDS",
"MAX_ASSERTION_BYTES",
"MAX_RESPONSE_BYTES",
"AnthropicIdentitySourceConfig",
"AnthropicIdentitySourceKind",
"AssertionReader",
"AssertionSource",
"AssertionSourceError",
"BodyEncoding",
"ExchangeError",
"ExchangeResult",
"InsecureTokenUrl",
"InternalIssuerSource",
"JwtBearerTokenExchangeEngine",
"KeycloakSource",
"MalformedTokenResponse",
"MintedToken",
"SecretReader",
"SigningKeyReader",
"SyncTokenPoster",
"TokenEndpointError",
"TokenExchangeSpec",
"TokenTransportError",
"build_jwk",
"build_jwks",
"default_token_exchange_engine",
"fetch_keycloak_assertion",
"identity_source_config_adapter",
"identity_source_ref",
"internal_issuer_assertion_source",
"internal_issuer_jwks_document",
"jwks_document_json",
"keycloak_assertion_source",
"load_es256_private_key",
"mint_internal_issuer_assertion",
"redact_oauth_error_body",
"rfc7638_thumbprint",
"sign_es256_jwt",
"validate_token_endpoint_url",
)

View file

@ -0,0 +1,225 @@
"""Fetches a fresh RFC 6749 client_credentials assertion for Anthropic's ``keycloak`` identity
source: LiteLLM authenticates to Keycloak as its own confidential client and presents the
resulting ``access_token`` as the workload assertion (Phase 1 decision 2).
The client secret is the operator-supplied pointer at ``KeycloakSource.client_secret_ref``,
resolved the same way every other WIF secret pointer already is (env, a Credential, or whatever
secret manager ``litellm.secret_manager_client`` is globally configured to, Vault included).
Every fetch is a fresh HTTP POST; nothing here caches a fetched token, since the outer
token-exchange engine already caches the Anthropic token it buys with one -- see decision 2's
"no Keycloak-side cache" ruling.
"""
import base64
import threading
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias
from urllib.parse import quote, quote_plus, urlencode
import httpx
from pydantic import BaseModel, SecretStr, ValidationError
from typing_extensions import assert_never
from litellm.llms.base_llm.auth.identity_source import KeycloakSource, ref_for_error_message
from litellm.llms.base_llm.auth.token_exchange import (
MAX_RESPONSE_BYTES,
endpoint_url_for_error_message,
redact_oauth_error_body,
require_posted_response,
validate_token_endpoint_url,
)
from litellm.llms.base_llm.auth.types import InsecureTokenUrl, SyncTokenPoster
if TYPE_CHECKING:
from litellm.llms.custom_httpx.http_handler import HTTPHandler
SecretReader: TypeAlias = Callable[[str], str | None]
_GRANT_TYPE: Final = "client_credentials"
_TIMEOUT_SECONDS: Final = 30.0
_FORM_CONTENT_TYPE: Final = "application/x-www-form-urlencoded"
class _ClientCredentialsResponse(BaseModel):
access_token: str
def _default_secret_reader(ref: str) -> str | None:
from litellm.secret_managers.main import get_secret_str
return get_secret_str(ref)
def _new_keycloak_handler() -> "HTTPHandler":
from litellm.llms.custom_httpx.http_handler import HTTPHandler
return HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0), follow_redirects=False)
class _HttpxSyncKeycloakPoster:
"""Dedicated HTTPHandler for the Keycloak token POST: no ``logging_obj`` (so litellm's
request/response logging never sees the client secret or the fetched token), redirects
disabled. A separate instance from the outer engine's own poster, since this is a genuinely
new HTTP call site whose no-logging guarantee must be built here, not assumed inherited."""
def __init__(self, handler_factory: Callable[[], "HTTPHandler"] = _new_keycloak_handler) -> None:
self._lock: Final = threading.Lock()
self._handler_factory: Final = handler_factory
self._handler: HTTPHandler | None = None
def _handler_instance(self) -> "HTTPHandler":
with self._lock:
if self._handler is None:
self._handler = self._handler_factory()
return self._handler
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
try:
response: Final[httpx.Response | None] = self._handler_instance().post( # pyright: ignore[reportUnknownMemberType] # HTTPHandler.post is legacy-untyped; the result is validated below
url,
content=content,
headers=dict(headers),
timeout=timeout,
)
except httpx.HTTPStatusError as e:
return e.response
return require_posted_response(response, "keycloak token endpoint")
_DEFAULT_POSTER: Final[SyncTokenPoster] = _HttpxSyncKeycloakPoster()
def _form_encode(value: str) -> str:
"""RFC 6749 Appendix B before RFC 6749 2.3.1's base64: application/x-www-form-urlencoded
with spaces as ``%20`` rather than ``+``, else a reserved character (":", "+", "%", " ") in
the id or secret corrupts the credential the far side decodes back out of Basic auth."""
return quote(value, safe="")
def _basic_auth_header(client_id: str, client_secret: str) -> str:
encoded_pair: Final = f"{_form_encode(client_id)}:{_form_encode(client_secret)}"
return "Basic " + base64.b64encode(encoded_pair.encode()).decode("ascii")
def _prepared_request(config: KeycloakSource, client_secret: str) -> tuple[bytes, Mapping[str, str]]:
scope_field: Final[Mapping[str, str]] = (
MappingProxyType({"scope": config.scope}) if config.scope else MappingProxyType({})
)
match config.auth_method:
case "client_secret_basic":
return (
urlencode(MappingProxyType({"grant_type": _GRANT_TYPE, **scope_field})).encode(),
MappingProxyType(
{
"content-type": _FORM_CONTENT_TYPE,
"authorization": _basic_auth_header(config.client_id, client_secret),
}
),
)
case "client_secret_post":
return (
urlencode(
MappingProxyType(
{
"grant_type": _GRANT_TYPE,
"client_id": config.client_id,
"client_secret": client_secret,
**scope_field,
}
)
).encode(),
MappingProxyType({"content-type": _FORM_CONTENT_TYPE}),
)
case _:
assert_never(config.auth_method)
def _resolve_client_secret(config: KeycloakSource, secret_reader: SecretReader) -> str:
secret: Final = secret_reader(config.client_secret_ref)
if not secret:
raise ValueError(f"keycloak client secret {ref_for_error_message(config.client_secret_ref)} could not be read")
return secret
def _wire_forms_of_secret(config: KeycloakSource, client_secret: str) -> tuple[SecretStr, ...]:
"""Every shape the secret leaves this process in, so an echo of any of them is caught.
Neither grant sends the secret verbatim. client_secret_basic base64s ``id:secret``, which
decodes straight back to it, and client_secret_post percent-escapes it. An endpoint echoing
either shape hands over reversible material a raw comparison would miss.
"""
raw: Final = SecretStr(client_secret)
match config.auth_method:
case "client_secret_basic":
encoded_pair: Final = f"{_form_encode(config.client_id)}:{_form_encode(client_secret)}"
return (raw, SecretStr(base64.b64encode(encoded_pair.encode()).decode("ascii")))
case "client_secret_post":
# urlencode escapes reserved characters and writes a space as "+", so a secret
# containing either leaves in a shape the raw comparison would not recognise coming
# back. quote_plus is what urlencode itself applies.
return (raw, SecretStr(quote_plus(client_secret)))
case _:
assert_never(config.auth_method)
def _endpoint_error_message(config: KeycloakSource, response: httpx.Response, client_secret: str) -> str:
endpoint_error: Final = redact_oauth_error_body(
response.status_code, response.text, _wire_forms_of_secret(config, client_secret)
)
return (
f"keycloak token endpoint {endpoint_url_for_error_message(config.token_url)} "
f"returned HTTP {endpoint_error.status_code}: {endpoint_error.redacted_body}"
)
def _parse_success_body(response: httpx.Response) -> str:
if len(response.content) > MAX_RESPONSE_BYTES:
raise ValueError("keycloak token response exceeded the size cap")
try:
parsed: Final = _ClientCredentialsResponse.model_validate_json(response.content)
except ValidationError as e:
raise ValueError("keycloak token response failed schema validation") from e
token: Final = parsed.access_token.strip()
if not token:
raise ValueError("keycloak token response carried an empty access_token")
return token
def fetch_keycloak_assertion(
config: KeycloakSource,
*,
poster: SyncTokenPoster = _DEFAULT_POSTER,
secret_reader: SecretReader = _default_secret_reader,
) -> str:
"""POSTs one fresh client_credentials grant and returns the resulting ``access_token`` as the
workload assertion; the caller must not cache the result -- see the module docstring."""
match validate_token_endpoint_url(config.token_url):
case InsecureTokenUrl(host=host):
raise ValueError(f"keycloak token_url must use https; refusing to send the client secret to host {host!r}")
case _:
pass
client_secret: Final = _resolve_client_secret(config, secret_reader)
content, headers = _prepared_request(config, client_secret)
try:
response: Final = poster.post(config.token_url, content=content, headers=headers, timeout=_TIMEOUT_SECONDS)
except Exception as e: # noqa: BLE001 # injected posters may raise beyond httpx; every failure becomes a ValueError
raise ValueError(
f"could not reach the keycloak token endpoint {endpoint_url_for_error_message(config.token_url)}: "
f"{type(e).__name__}"
) from e
if not 200 <= response.status_code < 300:
raise ValueError(_endpoint_error_message(config, response, client_secret))
return _parse_success_body(response)
def keycloak_assertion_source(
config: KeycloakSource,
*,
poster: SyncTokenPoster = _DEFAULT_POSTER,
secret_reader: SecretReader = _default_secret_reader,
) -> Callable[[], str]:
"""A zero-arg closure that fetches fresh on every call: the shape an ``oidc/keycloak/...``
ref dispatches to once wired into ``TokenExchangeSpec.assertion_source`` (Phase 1 decision 7)
-- the caller parses the config and closes this function over it, with no registry involved."""
return lambda: fetch_keycloak_assertion(config, poster=poster, secret_reader=secret_reader)

View file

@ -0,0 +1,76 @@
"""Tagged-union identity-source configs for Anthropic workload identity federation, beyond the
existing token_file/env resolver in ``litellm/llms/anthropic/wif.py``.
Each variant only ever carries secret *pointer names* (``signing_key_ref``, ``client_secret_ref``),
never a resolved secret value, so ``identity_source_ref`` can safely hash a variant into the short,
content-derived ``oidc/<kind>/<hash>`` string used elsewhere as a get_secret ref, a token-exchange
cache-key discriminator, and an operator-facing error pointer.
"""
import hashlib
from enum import Enum
from typing import Annotated, Final, Literal, TypeAlias
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
_REF_HASH_HEX_LENGTH: Final = 16
_MAX_TTL_SECONDS: Final = 3600
_DEFAULT_TTL_SECONDS: Final = 300
class AnthropicIdentitySourceKind(str, Enum):
internal_issuer = "internal_issuer"
keycloak = "keycloak"
class InternalIssuerSource(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", hide_input_in_errors=True)
kind: Literal[AnthropicIdentitySourceKind.internal_issuer] = AnthropicIdentitySourceKind.internal_issuer
issuer_url: str
subject: str
audience: str | None = None
ttl_seconds: Annotated[int, Field(gt=0, le=_MAX_TTL_SECONDS)] = _DEFAULT_TTL_SECONDS
signing_key_ref: str
class KeycloakSource(BaseModel):
model_config = ConfigDict(frozen=True, extra="forbid", hide_input_in_errors=True)
kind: Literal[AnthropicIdentitySourceKind.keycloak] = AnthropicIdentitySourceKind.keycloak
token_url: str
client_id: str
auth_method: Literal["client_secret_basic", "client_secret_post"] = "client_secret_basic"
client_secret_ref: str
scope: str | None = None
AnthropicIdentitySourceConfig: TypeAlias = Annotated[InternalIssuerSource | KeycloakSource, Field(discriminator="kind")]
identity_source_config_adapter: Final = TypeAdapter[AnthropicIdentitySourceConfig](AnthropicIdentitySourceConfig)
def identity_source_ref(config: AnthropicIdentitySourceConfig) -> str:
"""``oidc/<kind>/<hash>``: a short, secret-free pointer, stable for identical config and rolling
whenever any field does, including a ``*_ref`` pointer NAME (never the secret it points to)."""
digest: Final = hashlib.sha256(config.model_dump_json().encode()).hexdigest()[:_REF_HASH_HEX_LENGTH]
return f"oidc/{config.kind.value}/{digest}"
_POINTER_REF_PREFIXES: Final = (
"oidc/",
"os.environ/",
"hashicorp_vault/",
"aws_secret_manager/",
"google_secret_manager/",
)
def ref_for_error_message(ref: str) -> str:
"""A ``*_ref`` rendered for an operator-facing error.
Naming the pointer is deliberate: it is what tells an operator which setting failed to
resolve. But these fields only ever fail to resolve when what was written is not a pointer,
and an operator who pasted the secret itself has made the field's value the secret. So the
value is echoed only when it is recognizably a pointer, and withheld otherwise.
"""
return ref if ref.startswith(_POINTER_REF_PREFIXES) else "<withheld: not a secret reference>"

View file

@ -0,0 +1,86 @@
"""Mints a self-issued workload assertion for Anthropic's ``internal_issuer`` identity source:
LiteLLM signs its own short-lived ES256 JWT instead of reading one from a mounted OIDC file.
Signing custody is the operator-supplied PEM at ``InternalIssuerSource.signing_key_ref``,
resolved the same way every other WIF secret pointer already is (env, a Credential, or
whatever secret manager ``litellm.secret_manager_client`` is globally configured to, Vault
included) -- see Phase 1 decision 1. Every mint is fresh; nothing here caches a minted JWT,
since the outer token-exchange engine already caches the Anthropic token it buys with one.
"""
import time
import uuid
from collections.abc import Callable, Mapping
from types import MappingProxyType
from typing import Final, TypeAlias
from litellm.llms.base_llm.auth.identity_source import InternalIssuerSource, ref_for_error_message
from litellm.llms.base_llm.auth.jwt_signing import jwks_document_json, sign_es256_jwt
SigningKeyReader: TypeAlias = Callable[[str], str | None]
def _default_signing_key_reader(ref: str) -> str | None:
from litellm.secret_managers.main import get_secret_str
return get_secret_str(ref)
def _claims(config: InternalIssuerSource, issued_at: int) -> Mapping[str, object]:
return MappingProxyType(
{
key: value
for key, value in (
("sub", config.subject),
("iss", config.issuer_url),
("aud", config.audience),
("iat", issued_at),
("exp", issued_at + config.ttl_seconds),
("jti", str(uuid.uuid4())),
)
if value is not None
}
)
def _resolve_signing_key(config: InternalIssuerSource, key_reader: SigningKeyReader) -> str:
pem: Final = key_reader(config.signing_key_ref)
if not pem:
raise ValueError(
f"internal_issuer signing key {ref_for_error_message(config.signing_key_ref)} could not be read"
)
return pem
def mint_internal_issuer_assertion(
config: InternalIssuerSource,
*,
key_reader: SigningKeyReader = _default_signing_key_reader,
clock: Callable[[], float] = time.time,
) -> str:
"""Signs one fresh, short-lived assertion; the caller must not cache the result, since a
cached copy would defeat the point of re-minting on every exchange."""
pem: Final = _resolve_signing_key(config, key_reader)
return sign_es256_jwt(pem, _claims(config, issued_at=int(clock())))
def internal_issuer_assertion_source(
config: InternalIssuerSource,
*,
key_reader: SigningKeyReader = _default_signing_key_reader,
clock: Callable[[], float] = time.time,
) -> Callable[[], str]:
"""A zero-arg closure that mints fresh on every call: the shape an ``oidc/internal_issuer/...``
ref dispatches to once wired into ``TokenExchangeSpec.assertion_source`` (Phase 1 decision 7)
-- the caller parses the config and closes this function over it, with no registry involved."""
return lambda: mint_internal_issuer_assertion(config, key_reader=key_reader, clock=clock)
def internal_issuer_jwks_document(
config: InternalIssuerSource,
*,
key_reader: SigningKeyReader = _default_signing_key_reader,
) -> str:
"""The operator-facing JWKS export, resolved from a configured identity source rather than
a raw PEM in hand -- the JSON document to register as Anthropic's inline federation issuer."""
return jwks_document_json(_resolve_signing_key(config, key_reader))

View file

@ -0,0 +1,115 @@
"""ES256 JWT signing primitives for Anthropic workload identity federation's
``internal_issuer`` identity source (see ``identity_source.InternalIssuerSource``).
Pure functions over an already-resolved PEM string: no I/O, no secret-manager awareness, no
caching. Given the signing key at, say, $ISSUER_SIGNING_KEY_PEM, an operator publishes the
JWKS document Anthropic's inline federation issuer needs with one line:
python -c "from litellm.llms.base_llm.auth.jwt_signing import jwks_document_json; \\
import os; print(jwks_document_json(os.environ['ISSUER_SIGNING_KEY_PEM']))"
"""
import base64
import hashlib
import json
from collections.abc import Mapping
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, TypeAlias
if TYPE_CHECKING:
from cryptography.hazmat.primitives.asymmetric import ec
ALG: Final = "ES256"
MISSING_SIGNING_DEPENDENCIES_MESSAGE: Final = (
"the internal_issuer identity source needs PyJWT and cryptography, which a base litellm install "
"does not include: pip install 'litellm[proxy]'"
)
_JWK_CURVE_NAME: Final = "P-256"
_JWK_KEY_TYPE: Final = "EC"
_COORDINATE_BYTE_LENGTH: Final = 32 # P-256 field element width, RFC 7518 6.2.1.2/6.2.1.3
Jwk: TypeAlias = Mapping[str, str]
Jwks: TypeAlias = Mapping[str, tuple[Jwk, ...]]
def load_es256_private_key(pem: str) -> "ec.EllipticCurvePrivateKey":
"""Parses an unencrypted PEM EC private key. Never echoes the key material in an error."""
try:
from cryptography.hazmat.primitives.asymmetric import ec
from cryptography.hazmat.primitives.serialization import load_pem_private_key
except ImportError as e:
raise ImportError(MISSING_SIGNING_DEPENDENCIES_MESSAGE) from e
try:
key: Final = load_pem_private_key(pem.encode(), password=None)
except (ValueError, TypeError) as e:
raise ValueError("internal_issuer signing key is not a valid unencrypted PEM private key") from e
if not isinstance(key, ec.EllipticCurvePrivateKey) or not isinstance(key.curve, ec.SECP256R1):
raise ValueError( # noqa: TRY004 # the reader classifies ValueError into a readable config error; TypeError would not
"internal_issuer signing key must be an EC P-256 (secp256r1) private key for ES256"
)
return key
def _b64url_coordinate(value: int) -> str:
return base64.urlsafe_b64encode(value.to_bytes(_COORDINATE_BYTE_LENGTH, "big")).rstrip(b"=").decode("ascii")
def _jwk_thumbprint_members(public_key: "ec.EllipticCurvePublicKey") -> Jwk:
"""RFC 7638 3.2's exact EC member set (crv, kty, x, y) and nothing else: an extra member
here would change the thumbprint and desync it from the ``kid`` published in the JWKS."""
numbers: Final = public_key.public_numbers()
return MappingProxyType(
{
"crv": _JWK_CURVE_NAME,
"kty": _JWK_KEY_TYPE,
"x": _b64url_coordinate(numbers.x),
"y": _b64url_coordinate(numbers.y),
}
)
def rfc7638_thumbprint(public_key: "ec.EllipticCurvePublicKey") -> str:
"""RFC 7638: SHA-256 over the lexicographically member-ordered, whitespace-free JSON
rendering of the thumbprint members, base64url-encoded without padding."""
canonical: Final = json.dumps(
dict(sorted(_jwk_thumbprint_members(public_key).items())),
separators=(",", ":"),
)
return base64.urlsafe_b64encode(hashlib.sha256(canonical.encode()).digest()).rstrip(b"=").decode("ascii")
def build_jwk(public_key: "ec.EllipticCurvePublicKey", kid: str) -> Jwk:
return MappingProxyType({**_jwk_thumbprint_members(public_key), "use": "sig", "alg": ALG, "kid": kid})
def build_jwks(public_key: "ec.EllipticCurvePublicKey") -> Jwks:
kid: Final = rfc7638_thumbprint(public_key)
return MappingProxyType({"keys": (build_jwk(public_key, kid),)})
def jwks_document_json(pem: str) -> str:
"""The operator-facing export: the JSON document to register as Anthropic's inline JWKS.
``build_jwks`` returns ``MappingProxyType``/tuple values per this repo's no-mutation
convention; the ``json`` module only knows plain ``dict``/``list``, so those are converted
at this one serialization boundary rather than giving up immutability throughout the module.
"""
key: Final = load_es256_private_key(pem)
jwks: Final = build_jwks(key.public_key())
return json.dumps(
{"keys": [dict(jwk) for jwk in jwks["keys"]]},
indent=2,
)
def sign_es256_jwt(pem: str, claims: Mapping[str, object]) -> str:
"""Signs ``claims`` with the PEM key, stamping ``kid`` as its RFC 7638 thumbprint so a
verifier can look the signing key up in the published JWKS by ``kid`` alone."""
try:
import jwt
except ImportError as e:
raise ImportError(MISSING_SIGNING_DEPENDENCIES_MESSAGE) from e
key: Final = load_es256_private_key(pem)
kid: Final = rfc7638_thumbprint(key.public_key())
headers: Final = {"kid": kid}
return jwt.encode(dict(claims), key, algorithm=ALG, headers=headers)

View file

@ -0,0 +1,181 @@
"""Same-host token store for the JWT-bearer exchange engine.
Anthropic accepts an assertion carrying a ``jti`` once per issuer, so every uvicorn worker that
reads the same projected token file must share the token the first exchange minted instead of
re-sending the same assertion. The engine keys the store by cache key and only reuses a stored
token minted from the assertion it currently holds; a rotated assertion always buys a fresh token.
"""
import contextlib
import os
import sys
import tempfile
import threading
from collections.abc import Generator
from dataclasses import dataclass
from pathlib import Path
from typing import Final, Protocol
from pydantic import BaseModel, SecretStr, ValidationError
from litellm._logging import verbose_logger
CACHE_DIR_ENV: Final = "LITELLM_TOKEN_EXCHANGE_CACHE_DIR"
@dataclass(frozen=True, slots=True)
class StoredToken:
access_token: SecretStr
expires_at_epoch: float | None
assertion_sha256: str
class SharedTokenStore(Protocol):
"""Every method is best-effort: a store that cannot read, write, or lock degrades to a per-process
cache and never raises into the mint path."""
def load(self, key: str) -> StoredToken | None: ...
def save(self, key: str, token: StoredToken) -> None: ...
def delete(self, key: str) -> None: ...
def lock(self, key: str) -> contextlib.AbstractContextManager[None]: ...
class _StoredTokenFile(BaseModel):
access_token: str
expires_at_epoch: float | None
assertion_sha256: str
def _directory_is_private(directory: Path) -> bool:
try:
directory.mkdir(mode=0o700, exist_ok=True)
stat: Final = directory.stat()
except OSError as e:
verbose_logger.warning("Token exchange cache directory %s is unusable (%s); caching per process", directory, e)
return False
if stat.st_uid != os.getuid() or stat.st_mode & 0o077:
verbose_logger.warning(
"Token exchange cache directory %s must be owned by uid %d with mode 0700; caching per process",
directory,
os.getuid(),
)
return False
return True
def _unlink(path: Path) -> None:
with contextlib.suppress(OSError):
path.unlink()
def _write_token_file(directory: Path, key: str, body: bytes) -> None:
"""The token is staged in its own file and renamed over the entry, so a reader never sees a
half-written one. Every failure unlinks the staging file, including the buffered write that only
reaches the disk when the handle closes: nothing else sweeps this directory, and that file holds
a token that still works. The rename leaves nothing behind for the unlink to find."""
descriptor, name = tempfile.mkstemp(dir=directory, prefix=f"{key}.")
os.close(descriptor)
staged: Final = Path(name)
try:
staged.write_bytes(body)
os.replace(staged, directory / f"{key}.json")
finally:
_unlink(staged)
class FileTokenStore:
"""One ``<cache key>.json`` (mode 0600) and one ``<cache key>.lock`` (flock) per identity under a
directory only the proxy's uid can enter; the directory is checked on first use, not at import."""
def __init__(self, directory: Path) -> None:
self._directory: Final = directory
self._ready_lock: Final = threading.Lock()
self._ready: bool | None = None
@property
def directory(self) -> Path:
return self._directory
def _usable(self) -> bool:
with self._ready_lock:
if self._ready is None:
self._ready = _directory_is_private(self._directory)
return self._ready
def load(self, key: str) -> StoredToken | None:
if not self._usable():
return None
try:
raw: Final = (self._directory / f"{key}.json").read_bytes()
parsed: Final = _StoredTokenFile.model_validate_json(raw)
except FileNotFoundError:
return None
except (OSError, ValidationError) as e:
verbose_logger.debug("Ignoring unreadable token exchange cache entry: %s", e)
return None
return StoredToken(
access_token=SecretStr(parsed.access_token),
expires_at_epoch=parsed.expires_at_epoch,
assertion_sha256=parsed.assertion_sha256,
)
def save(self, key: str, token: StoredToken) -> None:
if not self._usable():
return
body: Final = (
_StoredTokenFile(
access_token=token.access_token.get_secret_value(),
expires_at_epoch=token.expires_at_epoch,
assertion_sha256=token.assertion_sha256,
)
.model_dump_json()
.encode()
)
try:
_write_token_file(self._directory, key, body)
except OSError as e:
verbose_logger.debug("Token exchange cache entry not written: %s", e)
def delete(self, key: str) -> None:
if not self._usable():
return
with contextlib.suppress(FileNotFoundError, OSError):
(self._directory / f"{key}.json").unlink()
@contextlib.contextmanager
def lock(self, key: str) -> Generator[None]:
if sys.platform == "win32" or not self._usable():
yield
return
import fcntl
try:
fd: Final = os.open(self._directory / f"{key}.lock", os.O_RDWR | os.O_CREAT, 0o600)
except OSError as e:
verbose_logger.debug("Token exchange cache lock unavailable (%s); minting without it", e)
yield
return
try:
fcntl.flock(fd, fcntl.LOCK_EX)
yield
finally:
with contextlib.suppress(OSError):
fcntl.flock(fd, fcntl.LOCK_UN)
os.close(fd)
def default_shared_token_store() -> SharedTokenStore | None:
"""``LITELLM_TOKEN_EXCHANGE_CACHE_DIR`` relocates the store; setting it empty disables it. Without
it the store lives under the temp directory, keyed by uid, so the workers of one proxy share it and
other users on the host cannot read it. Windows has no ``flock``, so it caches per process there."""
if sys.platform == "win32":
return None
configured: Final = os.environ.get(CACHE_DIR_ENV)
if configured == "":
return None
if configured is not None:
return FileTokenStore(Path(configured))
return FileTokenStore(Path(tempfile.gettempdir()) / f"litellm-token-exchange-{os.getuid()}")

View file

@ -0,0 +1,941 @@
"""RFC 7523 JWT-bearer token exchange engine, shared across providers.
One sync state machine per process: bounded engine-owned entry map, two-tier
refresh (advisory background refresh + mandatory single-flight), HTTPS pinning,
response caps, and RFC 6749 5.2 redaction. Providers describe a grant profile as
a ``TokenExchangeSpec`` and map the typed ``ExchangeError`` union to their own
public exception contract.
"""
import asyncio
import hashlib
import json
import re
import threading
import time
from collections.abc import Callable, Coroutine, Iterator, Mapping, Sequence
from concurrent.futures import Executor, ThreadPoolExecutor
from dataclasses import dataclass
from itertools import chain
from math import inf
from types import MappingProxyType
from typing import TYPE_CHECKING, Final, Protocol, TypeAlias
from urllib.parse import unquote, unquote_plus, urlencode, urlsplit, urlunsplit
import httpx
from pydantic import BaseModel, SecretStr, TypeAdapter, ValidationError
from typing_extensions import assert_never
from litellm._logging import verbose_logger
from litellm.llms.base_llm.auth.shared_token_store import SharedTokenStore, StoredToken, default_shared_token_store
from litellm.llms.base_llm.auth.types import (
AssertionReader,
AssertionSource,
AssertionSourceError,
ExchangeCallType,
ExchangeError,
ExchangeResult,
InsecureTokenUrl,
MalformedTokenResponse,
MintedToken,
SyncTokenPoster,
TokenEndpointError,
TokenExchangeMetricsSink,
TokenExchangeSpec,
TokenTransportError,
)
from litellm.types.services import ServiceTypes
if TYPE_CHECKING:
from litellm.llms.custom_httpx.http_handler import HTTPHandler
CALL_TYPE_COLD_MINT: Final[ExchangeCallType] = "cold_mint"
CALL_TYPE_MANDATORY_REFRESH: Final[ExchangeCallType] = "mandatory_refresh"
CALL_TYPE_ADVISORY_REFRESH: Final[ExchangeCallType] = "advisory_refresh"
CALL_TYPE_CACHE_HIT: Final = "cache_hit"
ADVISORY_REFRESH_SECONDS: Final = 120.0
MANDATORY_REFRESH_SECONDS: Final = 30.0
ADVISORY_REFRESH_LIFETIME_FRACTION: Final = 0.5
MANDATORY_REFRESH_LIFETIME_FRACTION: Final = 0.125
ADVISORY_REFRESH_BACKOFF_SECONDS: Final = 5.0
FALLBACK_TOKEN_TTL_SECONDS: Final = 60.0
# Metrics are best-effort, so the backlog is capped and further events are dropped. Request volume
# must not be able to grow this queue without bound when a telemetry backend stalls.
_METRICS_QUEUE_LIMIT: Final = 1000
MAX_ASSERTION_BYTES: Final = 16 * 1024
MAX_RESPONSE_BYTES: Final = 1024 * 1024
_REDACTION_CAP: Final = 256
_FOLLOWER_WAIT_GRACE_SECONDS: Final = 5.0
_LOCAL_HOSTS: Final = frozenset({"localhost", "127.0.0.1", "::1"})
_OAUTH_ERROR_FIELDS: Final = ("error", "error_description", "error_uri")
_NESTED_ERROR_FIELDS: Final = ("type", "message")
_CONTENT_TYPES: Final = MappingProxyType({"json": "application/json", "form": "application/x-www-form-urlencoded"})
_OVERSIZED_BODY_MESSAGE: Final = "oversized error response omitted"
_NON_OBJECT_BODY_MESSAGE: Final = "non-object error response omitted"
_NO_OAUTH_FIELDS_MESSAGE: Final = "error response carried no RFC 6749 fields"
_UNSTRUCTURED_BODY_MESSAGE: Final = "non-JSON error response omitted"
_REFLECTED_VALUE_MESSAGE: Final = "<redacted: response echoed the request>"
# A credential fragment shorter than this is not worth the false positives; longer, and a run
# shared with the assertion is reflection rather than coincidence.
_REFLECTION_MIN_RUN: Final = 8
# Everything a base64url credential is NOT made of, stripped so a fragment split by delimiters
# still lines up against the assertion.
_CREDENTIAL_CHARS: Final = re.compile(r"[^A-Za-z0-9._~+/=-]")
_SENTINEL_BODY_MESSAGES: Final = frozenset({_OVERSIZED_BODY_MESSAGE, _NON_OBJECT_BODY_MESSAGE})
class _TokenExchangeResponse(BaseModel):
access_token: str
expires_in: int | None = None
token_type: str | None = None
_RedactableBody: TypeAlias = Mapping[str, object] | list[object] | str | int | float | bool | None
_REDACTABLE_BODY_ADAPTER: Final = TypeAdapter[_RedactableBody](_RedactableBody)
def endpoint_url_for_error_message(url: str) -> str:
"""``url`` reduced to scheme, host and path for operator-facing errors.
A token endpoint is configuration, not a secret, and naming it is what makes these errors
actionable. But nothing stops an operator writing a credential into it, as a query parameter
or as userinfo, and these errors reach model callers, so neither part is echoed.
"""
parsed: Final = urlsplit(url)
host: Final = parsed.hostname or ""
authority: Final = f"{host}:{parsed.port}" if parsed.port is not None else host
return urlunsplit((parsed.scheme, authority, parsed.path, "", ""))
def validate_token_endpoint_url(url: str) -> str | InsecureTokenUrl:
parsed: Final = urlsplit(url)
if parsed.scheme == "https":
return url
if parsed.scheme == "http" and (parsed.hostname or "") in _LOCAL_HOSTS:
return url
return InsecureTokenUrl(host=parsed.hostname or "")
def redact_oauth_error_body(
status_code: int,
body_text: str,
assertion: SecretStr | Sequence[SecretStr] | None = None,
) -> TokenEndpointError:
"""``assertion`` may be every form of the credential that went out on the wire.
A grant that encodes its credential before sending it (``client_secret_basic`` base64s
``id:secret``) can have that encoded form echoed back, and it decodes straight to the secret,
so checking only the raw value lets reversible material through.
"""
rendered: Final = _redact_body_text(body_text)
secrets: Final = () if assertion is None else (assertion,) if isinstance(assertion, SecretStr) else tuple(assertion)
redacted: Final = next(
(
_REFLECTED_VALUE_MESSAGE
for secret in secrets
if _drop_reflected_assertion(rendered, secret) is _REFLECTED_VALUE_MESSAGE
),
rendered,
)
return TokenEndpointError(status_code=status_code, redacted_body=redacted)
def _drop_reflected_assertion(rendered: str, assertion: SecretStr | None) -> str:
"""Catches an endpoint that echoes the submitted credential back, verbatim or in fragments,
however it split or percent-encoded it.
Both sides are reduced to the characters a credential is made of before comparison. Stripping
only the rendered side would stop matching a secret that carries spaces or punctuation of its
own, which is exactly the hand-set passphrase most at risk of being echoed.
This stops an accidental or naive echo. It cannot stop an endpoint that deliberately re-encodes
or interleaves the credential, and it is not what keeps the credential from the endpoint, which
already holds it. What it protects is blast radius: keeping the value out of the caller's error
and out of third-party log sinks.
"""
if assertion is None:
return rendered
secret: Final = assertion.get_secret_value()
if not secret:
return rendered
if secret in rendered:
return _REFLECTED_VALUE_MESSAGE
compacted_secret: Final = _CREDENTIAL_CHARS.sub("", secret)
if not compacted_secret:
return rendered
return _REFLECTED_VALUE_MESSAGE if _shares_a_credential_run(rendered, compacted_secret) else rendered
def _shares_a_credential_run(rendered: str, compacted_secret: str) -> bool:
"""``unquote`` covers a credential sent form-encoded, without every caller enumerating that
shape for itself: percent-escaping is reversible and applies to any field, query string
included.
A secret shorter than the probe run is compared whole: a window longer than the secret can
never be found inside it, which would leave a short client secret unprotected in every shape
but the verbatim one.
"""
# unquote covers %XX; unquote_plus additionally covers the "+" a form-encoded body uses for a
# space. Both are kept rather than only the wider one, because "+" is a base64 character and
# decoding it away would lose a run that the undecoded candidate still matches on.
run: Final = min(_REFLECTION_MIN_RUN, len(compacted_secret))
compacted_candidates: Final = tuple(
_CREDENTIAL_CHARS.sub("", candidate) for candidate in (rendered, unquote(rendered), unquote_plus(rendered))
)
windows: Final = chain.from_iterable(_character_runs(candidate, run) for candidate in compacted_candidates)
return any(window in compacted_secret for window in windows)
def _character_runs(compacted: str, run: int) -> Iterator[str]:
return (compacted[start : start + run] for start in range(len(compacted) - run + 1))
def _redact_body_text(body_text: str) -> str:
if body_text in _SENTINEL_BODY_MESSAGES:
return body_text
if len(body_text) > MAX_RESPONSE_BYTES:
return _OVERSIZED_BODY_MESSAGE
try:
parsed: Final = _REDACTABLE_BODY_ADAPTER.validate_json(body_text)
except ValidationError:
return _UNSTRUCTURED_BODY_MESSAGE
match parsed:
case Mapping():
return _format_oauth_error_fields(parsed)
case _:
return _NON_OBJECT_BODY_MESSAGE
def _format_oauth_error_fields(body: Mapping[str, object]) -> str:
fields: Final = tuple(
f"{name}: {_format_oauth_error_value(body[name])}" for name in _OAUTH_ERROR_FIELDS if body.get(name) is not None
)
return "; ".join(fields) if fields else _NO_OAUTH_FIELDS_MESSAGE
def _format_oauth_error_value(value: object) -> str:
"""RFC 6749 types ``error`` as a string, but Anthropic (and other providers) nest their
own ``{"type": ..., "message": ...}`` envelope there; render that rather than a dict repr."""
if isinstance(value, Mapping):
nested: Final = tuple(
str(value[key])[:_REDACTION_CAP] for key in _NESTED_ERROR_FIELDS if value.get(key) is not None
)
if nested:
return " - ".join(nested)
return str(value)[:_REDACTION_CAP]
def _error_summary(error: ExchangeError) -> str:
match error:
case AssertionSourceError():
return f"AssertionSourceError: assertion {error.kind} from {error.source_ref}"
case InsecureTokenUrl():
return f"InsecureTokenUrl: insecure token endpoint host {error.host}"
case TokenEndpointError():
return f"TokenEndpointError: HTTP {error.status_code}: {error.redacted_body}"
case TokenTransportError():
return f"TokenTransportError: {error.detail}"
case MalformedTokenResponse():
return f"MalformedTokenResponse: {error.detail}"
case _:
assert_never(error)
class _MetricsFailure(Exception):
"""Never raised: typed carriers handed to the service failure hook so the prometheus
``error_class`` label names the ``ExchangeError`` variant; the message is the redacted
``_error_summary`` and carries no credential material."""
class TokenExchangeAssertionSourceFailure(_MetricsFailure): ...
class TokenExchangeInsecureUrlFailure(_MetricsFailure): ...
class TokenExchangeEndpointFailure(_MetricsFailure): ...
class TokenExchangeTransportFailure(_MetricsFailure): ...
class TokenExchangeMalformedResponseFailure(_MetricsFailure): ...
def _failure_exception(error: ExchangeError) -> _MetricsFailure:
summary: Final = _error_summary(error)
match error:
case AssertionSourceError():
return TokenExchangeAssertionSourceFailure(summary)
case InsecureTokenUrl():
return TokenExchangeInsecureUrlFailure(summary)
case TokenEndpointError():
return TokenExchangeEndpointFailure(summary)
case TokenTransportError():
return TokenExchangeTransportFailure(summary)
case MalformedTokenResponse():
return TokenExchangeMalformedResponseFailure(summary)
case _:
assert_never(error)
def _cache_key(spec: TokenExchangeSpec) -> str:
return hashlib.sha256(
"\x1f".join((spec.token_url, spec.assertion_ref, *spec.cache_key_identity)).encode()
).hexdigest()
def _shares_one_assertion_across_workers(spec: TokenExchangeSpec) -> bool:
"""The store exists so the workers reading one projected token file don't each spend that file's
single-use ``jti``. A source that mints its own assertion per exchange shares nothing with another
worker, so it never reads the store, never finds a hit there, and keeps its minted token off disk."""
return spec.assertion_source is None
def _assertion_digest(assertion: SecretStr) -> str:
return hashlib.sha256(assertion.get_secret_value().encode()).hexdigest()
def _assertion_fetch(reader: AssertionReader, spec: TokenExchangeSpec) -> AssertionSource:
"""``spec.assertion_source`` (an identity source's own fetch/mint closure) takes priority over
the engine-level reader when set; either way, failures are reported against ``spec.assertion_ref``."""
if spec.assertion_source is not None:
return spec.assertion_source
return lambda: reader(spec.assertion_ref)
def _read_assertion(fetch: AssertionSource, ref: str) -> SecretStr | AssertionSourceError:
from litellm.secret_managers.main import OidcPathNotAllowedError
try:
raw: Final = fetch()
except OidcPathNotAllowedError:
return AssertionSourceError(kind="disallowed_path", source_ref=ref)
except (ValueError, ImportError) as e:
return AssertionSourceError(kind="unreadable", source_ref=ref, detail=str(e)[:_REDACTION_CAP])
except Exception: # noqa: BLE001 # injected readers (secret managers) raise arbitrarily; all failures become values
return AssertionSourceError(kind="unreadable", source_ref=ref)
if raw is None:
return AssertionSourceError(kind="missing", source_ref=ref)
stripped: Final = raw.strip()
if not stripped:
return AssertionSourceError(kind="empty", source_ref=ref)
if len(stripped.encode("utf-8")) > MAX_ASSERTION_BYTES:
return AssertionSourceError(kind="oversized", source_ref=ref)
return SecretStr(stripped)
def _serialize_body(spec: TokenExchangeSpec, assertion: SecretStr) -> bytes:
if spec.body_encoding == "json":
return json.dumps(
{
**spec.static_body,
spec.assertion_field: assertion.get_secret_value(),
}
).encode()
return urlencode(
{
**spec.static_body,
spec.assertion_field: assertion.get_secret_value(),
}
).encode()
def _sanitize_expires_in(expires_in: int | None) -> float:
if expires_in is None or expires_in <= 0:
return FALLBACK_TOKEN_TTL_SECONDS
return float(expires_in)
@dataclass(frozen=True, slots=True)
class _RefreshWindows:
advisory: float
mandatory: float
def _refresh_windows(lifetime_seconds: float | None) -> _RefreshWindows:
"""A token whose whole life is shorter than the flat windows sits inside them from the moment it
is minted, so every request would arm another background exchange against the token endpoint.
Scaling each window by a fraction of the observed lifetime makes a 60s token refresh around its
half life instead; at a lifetime of 240s and above both fractions reach the flat windows, so
ordinary long-lived tokens keep exactly the 120s/30s behaviour."""
if lifetime_seconds is None or lifetime_seconds <= 0.0:
return _RefreshWindows(advisory=ADVISORY_REFRESH_SECONDS, mandatory=MANDATORY_REFRESH_SECONDS)
return _RefreshWindows(
advisory=min(ADVISORY_REFRESH_SECONDS, lifetime_seconds * ADVISORY_REFRESH_LIFETIME_FRACTION),
mandatory=min(MANDATORY_REFRESH_SECONDS, lifetime_seconds * MANDATORY_REFRESH_LIFETIME_FRACTION),
)
def _capped_body_text(response: httpx.Response) -> str:
if len(response.content) > MAX_RESPONSE_BYTES:
return _OVERSIZED_BODY_MESSAGE
return response.text
def _default_assertion_reader(ref: str) -> str | None:
from litellm.secret_managers.main import get_secret_str
return get_secret_str(ref)
def _new_exchange_handler() -> "HTTPHandler":
from litellm.llms.custom_httpx.http_handler import HTTPHandler
return HTTPHandler(timeout=httpx.Timeout(timeout=30.0, connect=5.0), follow_redirects=False)
def require_posted_response(response: httpx.Response | None, endpoint_label: str) -> httpx.Response:
"""The legacy ``HTTPHandler`` carries no return annotation, so a patched or stubbed client can
hand a poster ``None`` back; a transport error beats dereferencing it."""
if response is None:
raise httpx.TransportError(f"{endpoint_label} returned no response")
return response
class _HttpxSyncTokenPoster:
"""Default poster: a dedicated HTTPHandler (no logging_obj, so litellm's
pre/post-call body logging never sees the exchange POST); returns the
response for any status."""
def __init__(self, handler_factory: Callable[[], "HTTPHandler"] = _new_exchange_handler) -> None:
self._lock: Final = threading.Lock()
self._handler_factory: Final = handler_factory
self._handler: HTTPHandler | None = None
def _handler_instance(self) -> "HTTPHandler":
with self._lock:
if self._handler is None:
self._handler = self._handler_factory()
return self._handler
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response:
try:
response: Final[httpx.Response | None] = self._handler_instance().post( # pyright: ignore[reportUnknownMemberType] # HTTPHandler.post is legacy-untyped; the result is validated below
url,
content=content,
headers=dict(headers),
timeout=timeout,
)
except httpx.HTTPStatusError as e:
return e.response
return require_posted_response(response, "token endpoint")
class _ServiceLoggingHooks(Protocol):
"""The slice of ``litellm._service_logger.ServiceLogging`` the metrics sink calls; a protocol
so tests inject a recorder instead of monkeypatching."""
async def async_service_success_hook(self, service: ServiceTypes, call_type: str, duration: float) -> None: ...
async def async_service_failure_hook(
self, service: ServiceTypes, duration: float, error: str | Exception, call_type: str
) -> None: ...
_HooksCoroFactory: TypeAlias = Callable[
[_ServiceLoggingHooks],
Coroutine[object, object, None],
]
def _default_service_logging() -> _ServiceLoggingHooks:
from litellm._service_logger import ServiceLogging
return ServiceLogging()
class ServiceLoggingMetricsSink:
"""Default sink: bridges engine metrics onto litellm's ServiceTypes pattern
(prometheus ``litellm_anthropic_wif_*`` via ``service_callback``). The engine's entry points
are sync threads with no event loop, and the service hooks are async, so every emission is
fire-and-forget on a dedicated single worker thread that owns its own short-lived loop --
the mint path only ever pays for an executor queue put."""
def __init__(
self,
service_logging_factory: Callable[[], _ServiceLoggingHooks] = _default_service_logging,
executor: Executor | None = None,
) -> None:
self._lock: Final = threading.Lock()
self._service_logging_factory: Final = service_logging_factory
self._service_logging: _ServiceLoggingHooks | None = None
self._executor: Executor | None = executor
self._queued: int = 0
def _service_logging_instance(self) -> _ServiceLoggingHooks:
with self._lock:
if self._service_logging is None:
self._service_logging = self._service_logging_factory()
return self._service_logging
def _executor_instance(self) -> Executor:
with self._lock:
if self._executor is None:
self._executor = ThreadPoolExecutor(max_workers=1, thread_name_prefix="litellm-token-exchange-metrics")
return self._executor
def _emit(self, coro_factory: _HooksCoroFactory) -> None:
try:
asyncio.run(coro_factory(self._service_logging_instance()))
except Exception as e: # noqa: BLE001 # metrics are best-effort; emission failures must never surface
verbose_logger.debug("token exchange metrics emission failed: %s", e)
def _submit(self, coro_factory: _HooksCoroFactory) -> None:
"""Drop the event rather than queue it once the backlog is full. A stalled telemetry
backend must not let request volume grow an unbounded queue in the proxy: losing a
metric sample is always cheaper than losing the process."""
with self._lock:
if self._queued >= _METRICS_QUEUE_LIMIT:
verbose_logger.debug("token exchange metrics queue full, dropping event")
return
self._queued += 1
try:
self._executor_instance().submit(self._emit_and_release, coro_factory)
except Exception as e: # noqa: BLE001 # a rejected submit must not surface to the mint
with self._lock:
self._queued -= 1
verbose_logger.debug("token exchange metrics submit failed: %s", e)
def _emit_and_release(self, coro_factory: _HooksCoroFactory) -> None:
try:
self._emit(coro_factory)
finally:
with self._lock:
self._queued -= 1
def exchange_success(self, *, call_type: ExchangeCallType, duration_seconds: float) -> None:
def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]:
return hooks.async_service_success_hook(
service=ServiceTypes.ANTHROPIC_WIF, call_type=call_type, duration=duration_seconds
)
self._submit(start)
def exchange_failure(self, *, call_type: ExchangeCallType, duration_seconds: float, error: ExchangeError) -> None:
failure: Final = _failure_exception(error)
def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]:
return hooks.async_service_failure_hook(
service=ServiceTypes.ANTHROPIC_WIF, duration=duration_seconds, error=failure, call_type=call_type
)
self._submit(start)
def cache_hit(self) -> None:
def start(hooks: _ServiceLoggingHooks) -> Coroutine[object, object, None]:
return hooks.async_service_success_hook(
service=ServiceTypes.ANTHROPIC_WIF_CACHE, call_type=CALL_TYPE_CACHE_HIT, duration=0.0
)
self._submit(start)
class _Entry:
"""Single-flight state for one cache key; mutable by design, confined to the
engine, and only ever mutated under the engine lock."""
__slots__ = ("backoff_until", "done", "force_refresh", "in_flight", "last_error", "lifetime_seconds", "token")
def __init__(self, force_refresh: bool = False) -> None:
self.token: MintedToken | None = None
self.lifetime_seconds: float | None = None
self.in_flight: bool = False
self.done: Final = threading.Event()
self.backoff_until: float = float("-inf")
self.force_refresh: bool = force_refresh
self.last_error: ExchangeError | None = None
def arm(self) -> None:
self.in_flight = True
self.last_error = None
self.done.clear()
def disarm(self, now: float) -> None:
"""Undo ``arm`` for a refresh that never started. Nothing is on its way to publish, so the
entry must stop reading as in-flight, and the backoff keeps every later caller from
re-attempting a schedule that just failed."""
self.backoff_until = now + ADVISORY_REFRESH_BACKOFF_SECONDS
self.in_flight = False
self.done.set()
def _store(self, token: MintedToken, now: float) -> None:
self.token = token
self.lifetime_seconds = None if token.expires_at is None else max(token.expires_at - now, 0.0)
self.last_error = None
def publish(self, result: ExchangeResult, now: float) -> None:
match result:
case MintedToken():
self._store(result, now)
case _:
self.last_error = result
self.backoff_until = now + ADVISORY_REFRESH_BACKOFF_SECONDS
self.force_refresh = False
self.in_flight = False
self.done.set()
def publish_advisory(self, result: ExchangeResult, now: float) -> None:
"""A failed advisory refresh records only the backoff, never ``last_error``: a follower whose
cached token expires while this runs must be free to re-lead a fresh mint and recover."""
match result:
case MintedToken():
self._store(result, now)
case _:
self.backoff_until = now + ADVISORY_REFRESH_BACKOFF_SECONDS
self.in_flight = False
self.done.set()
@dataclass(frozen=True, slots=True)
class _Serve:
token: MintedToken
@dataclass(frozen=True, slots=True)
class _ServeAndRefresh:
token: MintedToken
@dataclass(frozen=True, slots=True)
class _Lead:
call_type: ExchangeCallType
@dataclass(frozen=True, slots=True)
class _Follow:
pass
@dataclass(frozen=True, slots=True)
class _Fail:
error: ExchangeError
_Decision: TypeAlias = _Serve | _ServeAndRefresh | _Lead | _Follow | _Fail
@dataclass(frozen=True, slots=True)
class _Unauthorized:
response: httpx.Response
assertion: SecretStr
def _denied(attempt: _Unauthorized) -> TokenEndpointError:
return redact_oauth_error_body(attempt.response.status_code, _capped_body_text(attempt.response), attempt.assertion)
class JwtBearerTokenExchangeEngine:
def __init__(
self,
poster: SyncTokenPoster | None = None,
assertion_reader: AssertionReader | None = None,
clock: Callable[[], float] = time.monotonic,
refresh_executor: Executor | None = None,
max_entries: int = 64,
metrics_sink: TokenExchangeMetricsSink | None = None,
shared_store: SharedTokenStore | None = None,
wall_clock: Callable[[], float] = time.time,
) -> None:
self._poster: Final[SyncTokenPoster] = poster if poster is not None else _HttpxSyncTokenPoster()
self._assertion_reader: Final[AssertionReader] = (
assertion_reader if assertion_reader is not None else _default_assertion_reader
)
self._clock: Final = clock
self._refresh_executor: Executor | None = refresh_executor
self._max_entries: Final = max_entries
self._metrics_sink: Final[TokenExchangeMetricsSink] = (
metrics_sink if metrics_sink is not None else ServiceLoggingMetricsSink()
)
self._shared_store: Final = shared_store
self._wall_clock: Final = wall_clock
self._lock: Final = threading.Lock()
self._entries: Final[dict[str, _Entry]] = {} # mutable-ok: engine-owned map guarded by _lock
def get_token(self, spec: TokenExchangeSpec) -> ExchangeResult:
"""A follower whose leader published nothing re-classifies rather than recursing, so a
contended entry cannot grow the stack one frame per failed leader."""
while True:
with self._lock:
entry = self._get_or_create_entry_locked(spec)
decision = self._classify_and_arm_locked(entry)
match decision:
case _Serve(token=token):
self._report_cache_hit()
return token
case _ServeAndRefresh(token=token):
self._report_cache_hit()
self._submit_advisory_refresh(spec, entry)
return token
case _Fail(error=error):
return error
case _Lead(call_type=call_type):
return self._lead(spec, entry, call_type)
case _Follow():
followed = self._await_leader(spec, entry)
if followed is not None:
return followed
case _:
assert_never(decision)
async def aget_token(self, spec: TokenExchangeSpec) -> ExchangeResult:
return await asyncio.to_thread(self.get_token, spec)
def invalidate(self, spec: TokenExchangeSpec) -> None:
key: Final = _cache_key(spec)
with self._lock:
if key in self._entries:
self._entries[key] = _Entry(force_refresh=True)
if self._shared_store is not None:
self._shared_store.delete(key)
def _get_or_create_entry_locked(self, spec: TokenExchangeSpec) -> _Entry:
key: Final = _cache_key(spec)
existing: Final = self._entries.get(key)
if existing is not None:
return existing
if len(self._entries) >= self._max_entries:
self._evict_locked()
created: Final = _Entry()
self._entries[key] = created
return created
def _evict_locked(self) -> None:
now: Final = self._clock()
stale: Final = tuple(
key
for key, entry in self._entries.items()
if not entry.in_flight
and (entry.token is None or (entry.token.expires_at is not None and entry.token.expires_at <= now))
)
for key in stale:
del self._entries[key]
if len(self._entries) < self._max_entries:
return
# Evict soonest-to-expire first, and take as many as the overshoot needs rather than one, so a
# burst of distinct identities does not leave the map permanently above max_entries. An entry
# a leader owns or a follower waits on is never a candidate, so a moment where every entry is
# in flight still over-inserts; that residue is bounded by the concurrent mints themselves.
evictable: Final = sorted(
(
entry.token.expires_at if entry.token is not None and entry.token.expires_at is not None else -inf,
key,
)
for key, entry in self._entries.items()
if not entry.in_flight
)
for _, key in evictable[: len(self._entries) - self._max_entries + 1]:
del self._entries[key]
def _classify_and_arm_locked(self, entry: _Entry) -> _Decision:
token: Final = entry.token
if token is not None and not entry.force_refresh:
if token.expires_at is None:
return _Serve(token=token)
windows: Final = _refresh_windows(entry.lifetime_seconds)
remaining: Final = token.expires_at - self._clock()
if remaining > windows.advisory:
return _Serve(token=token)
if remaining > windows.mandatory:
if entry.in_flight or self._clock() < entry.backoff_until:
return _Serve(token=token)
entry.arm()
return _ServeAndRefresh(token=token)
if entry.in_flight:
return _Follow()
if entry.last_error is not None and self._clock() < entry.backoff_until:
return _Fail(error=entry.last_error)
entry.arm()
return _Lead(call_type=CALL_TYPE_COLD_MINT if token is None else CALL_TYPE_MANDATORY_REFRESH)
def _executor_instance(self) -> Executor:
with self._lock:
if self._refresh_executor is None:
self._refresh_executor = ThreadPoolExecutor(thread_name_prefix="litellm-token-exchange-refresh")
return self._refresh_executor
def _submit_advisory_refresh(self, spec: TokenExchangeSpec, entry: _Entry) -> None:
"""The entry is already armed, so an executor that refuses the work would leave it reading
as in-flight with nothing on its way to publish, and every later caller would wait out the
follower timeout and fail. A refused submit disarms it and the cached token keeps serving."""
try:
self._executor_instance().submit(self._advisory_refresh, spec, entry)
except RuntimeError as e:
verbose_logger.debug("token exchange advisory refresh could not be scheduled: %s", e)
with self._lock:
entry.disarm(self._clock())
def _lead(self, spec: TokenExchangeSpec, entry: _Entry, call_type: ExchangeCallType) -> ExchangeResult:
started: Final = self._clock()
result: Final = self._exchange_never_raises(spec)
duration: Final = self._clock() - started
with self._lock:
entry.publish(result, now=self._clock())
self._report_exchange(call_type, duration, result)
return result
def _await_leader(self, spec: TokenExchangeSpec, entry: _Entry) -> "ExchangeResult | None":
"""None means the finished round left neither a valid token nor an error
(a failed advisory refresh); the caller re-enters and leads a fresh exchange."""
leader_finished: Final = entry.done.wait(2 * spec.timeout_seconds + _FOLLOWER_WAIT_GRACE_SECONDS)
with self._lock:
token: Final = entry.token
if token is not None and (token.expires_at is None or token.expires_at > self._clock()):
return token
if entry.last_error is not None:
return entry.last_error
if leader_finished:
return None
return TokenTransportError(detail="timed out waiting for the token exchange leader")
def _advisory_refresh(self, spec: TokenExchangeSpec, entry: _Entry) -> None:
started: Final = self._clock()
result: Final = self._exchange_never_raises(spec)
duration: Final = self._clock() - started
with self._lock:
now: Final = self._clock()
entry.publish_advisory(result, now=now)
stale_expires_at: Final = entry.token.expires_at if entry.token is not None else None
stale_mandatory: Final = _refresh_windows(entry.lifetime_seconds).mandatory
self._report_exchange(CALL_TYPE_ADVISORY_REFRESH, duration, result)
if isinstance(result, MintedToken):
return
seconds_to_mandatory_wall: Final = (
max(stale_expires_at - now - stale_mandatory, 0.0) if stale_expires_at is not None else 0.0
)
verbose_logger.warning(
"Advisory token refresh against %s failed (%s); serving the cached token for up to "
"%.0fs before the mandatory refresh wall; next attempt after %.0fs backoff",
urlsplit(spec.token_url).hostname or "",
_error_summary(result),
seconds_to_mandatory_wall,
ADVISORY_REFRESH_BACKOFF_SECONDS,
)
def _report_exchange(self, call_type: ExchangeCallType, duration_seconds: float, result: ExchangeResult) -> None:
try:
match result:
case MintedToken():
self._metrics_sink.exchange_success(call_type=call_type, duration_seconds=duration_seconds)
case _:
self._metrics_sink.exchange_failure(
call_type=call_type, duration_seconds=duration_seconds, error=result
)
except Exception as e: # noqa: BLE001 # metrics are best-effort; a sink failure must never fail a mint
verbose_logger.debug("token exchange metrics emission failed: %s", e)
def _report_cache_hit(self) -> None:
try:
self._metrics_sink.cache_hit()
except Exception as e: # noqa: BLE001 # metrics are best-effort; a sink failure must never fail a serve
verbose_logger.debug("token exchange cache-hit metric emission failed: %s", e)
def _exchange_never_raises(self, spec: TokenExchangeSpec) -> ExchangeResult:
"""The single-flight leader and the advisory refresher must always publish a result: an
unhandled exception here would leave the entry armed (in_flight, cleared event) forever, so
every subsequent caller for this key would follow a leader that never finishes."""
try:
return self._exchange(spec)
except Exception as e: # noqa: BLE001 # a leader must resolve its entry; any failure becomes a value
return TokenTransportError(detail=f"{type(e).__name__}: {e}"[:_REDACTION_CAP])
def _exchange(self, spec: TokenExchangeSpec) -> ExchangeResult:
url_check: Final = validate_token_endpoint_url(spec.token_url)
if isinstance(url_check, InsecureTokenUrl):
return url_check
fetch: Final = _assertion_fetch(self._assertion_reader, spec)
assertion: Final = _read_assertion(fetch, spec.assertion_ref)
if isinstance(assertion, AssertionSourceError):
return assertion
if self._shared_store is None or not _shares_one_assertion_across_workers(spec):
return self._mint(spec, fetch, assertion)
key: Final = _cache_key(spec)
with self._shared_store.lock(key):
shared: Final = self._shared_token(self._shared_store.load(key), _assertion_digest(assertion))
if shared is not None:
return shared
minted: Final = self._mint(spec, fetch, assertion)
if isinstance(minted, MintedToken):
self._shared_store.save(key, self._stored_token(minted))
return minted
def _shared_token(self, stored: StoredToken | None, assertion_sha256: str) -> MintedToken | None:
"""A stored token minted from the very assertion this process holds is the token that assertion
bought: another worker sharing the token file already exchanged it, and an issuer enforcing
single-use ``jti`` would only deny a second exchange."""
if stored is None or stored.assertion_sha256 != assertion_sha256:
return None
if stored.expires_at_epoch is None:
return MintedToken(access_token=stored.access_token, expires_at=None, assertion_sha256=assertion_sha256)
remaining: Final = stored.expires_at_epoch - self._wall_clock()
if remaining <= 0.0:
return None
return MintedToken(
access_token=stored.access_token,
expires_at=self._clock() + remaining,
assertion_sha256=assertion_sha256,
)
def _stored_token(self, token: MintedToken) -> StoredToken:
return StoredToken(
access_token=token.access_token,
expires_at_epoch=(
None if token.expires_at is None else self._wall_clock() + (token.expires_at - self._clock())
),
assertion_sha256=token.assertion_sha256,
)
def _mint(self, spec: TokenExchangeSpec, fetch: AssertionSource, assertion: SecretStr) -> ExchangeResult:
"""One 401 earns one retry, and only with an assertion that changed since the first attempt: a
token file rotated between the read and the POST is worth resending, the same assertion is not,
since an issuer that already consumed its ``jti`` denies it again."""
first: Final = self._post_assertion(spec, assertion)
if not isinstance(first, _Unauthorized):
return first
reread: Final = _read_assertion(fetch, spec.assertion_ref)
if isinstance(reread, AssertionSourceError):
return reread
if reread.get_secret_value() == assertion.get_secret_value():
return _denied(first)
second: Final = self._post_assertion(spec, reread)
if isinstance(second, _Unauthorized):
return _denied(second)
return second
def _post_assertion(self, spec: TokenExchangeSpec, assertion: SecretStr) -> "ExchangeResult | _Unauthorized":
try:
response: Final = self._poster.post(
spec.token_url,
content=_serialize_body(spec, assertion),
headers=MappingProxyType({"content-type": _CONTENT_TYPES[spec.body_encoding], **spec.request_headers}),
timeout=spec.timeout_seconds,
)
except Exception as e: # noqa: BLE001 # injected posters may raise beyond httpx; transport failures become values
return TokenTransportError(detail=f"{type(e).__name__}: {e}"[:_REDACTION_CAP])
if response.status_code == 401:
return _Unauthorized(response=response, assertion=assertion)
return self._parse_response(response, assertion)
def _parse_response(self, response: httpx.Response, assertion: SecretStr) -> ExchangeResult:
if not 200 <= response.status_code < 300:
return redact_oauth_error_body(response.status_code, _capped_body_text(response), assertion)
if len(response.content) > MAX_RESPONSE_BYTES:
return MalformedTokenResponse(detail="token response body exceeds the 1 MiB cap")
try:
parsed: Final = _TokenExchangeResponse.model_validate_json(response.content)
except ValidationError:
return MalformedTokenResponse(detail="token response failed RFC 6749 5.1 schema validation")
if parsed.token_type is not None and parsed.token_type.lower() != "bearer":
return MalformedTokenResponse(detail="token response carried a non-bearer token_type")
if not parsed.access_token.strip():
return MalformedTokenResponse(detail="token response carried an empty access_token")
return MintedToken(
access_token=SecretStr(parsed.access_token),
expires_at=self._clock() + _sanitize_expires_in(parsed.expires_in),
assertion_sha256=_assertion_digest(assertion),
)
default_token_exchange_engine: Final = JwtBearerTokenExchangeEngine(shared_store=default_shared_token_store())

View file

@ -0,0 +1,100 @@
"""Provider-agnostic types for the RFC 7523 JWT-bearer token exchange engine."""
from collections.abc import Callable, Mapping
from dataclasses import dataclass
from typing import Literal, Protocol, TypeAlias
import httpx
from pydantic import SecretStr
BodyEncoding: TypeAlias = Literal["json", "form"]
AssertionReader: TypeAlias = Callable[[str], str | None]
AssertionSource: TypeAlias = Callable[[], str | None]
@dataclass(frozen=True, slots=True)
class TokenExchangeSpec:
"""One grant profile as pure data: one instance per (provider, deployment, identity).
``token_url`` must be derived from deployment config/env only, never per-request caller
input. ``assertion_ref`` is a ``oidc/...`` get_secret ref resolved fresh on every exchange.
``assertion_source``, when set, is a zero-arg per-config fetch/mint closure that the engine
prefers over its own engine-level ``AssertionReader`` -- the dispatch mechanism identity
sources beyond token_file/env (e.g. ``internal_issuer``, ``keycloak``) use to plug into the
shared engine without a global registry. ``assertion_ref`` still names the cache-key
discriminator and the ref echoed into operator-facing errors either way.
"""
token_url: str
assertion_ref: str
assertion_field: str
static_body: Mapping[str, str]
body_encoding: BodyEncoding
request_headers: Mapping[str, str]
cache_key_identity: tuple[str, ...]
timeout_seconds: float = 30.0
assertion_source: AssertionSource | None = None
@dataclass(frozen=True, slots=True)
class MintedToken:
access_token: SecretStr
expires_at: float | None
assertion_sha256: str
@dataclass(frozen=True, slots=True)
class AssertionSourceError:
kind: Literal["missing", "empty", "oversized", "unreadable", "disallowed_path"]
source_ref: str
detail: str | None = None
@dataclass(frozen=True, slots=True)
class InsecureTokenUrl:
host: str
@dataclass(frozen=True, slots=True)
class TokenEndpointError:
status_code: int
redacted_body: str
@dataclass(frozen=True, slots=True)
class TokenTransportError:
detail: str
@dataclass(frozen=True, slots=True)
class MalformedTokenResponse:
detail: str
ExchangeError: TypeAlias = (
AssertionSourceError | InsecureTokenUrl | TokenEndpointError | TokenTransportError | MalformedTokenResponse
)
ExchangeResult: TypeAlias = MintedToken | ExchangeError
ExchangeCallType: TypeAlias = Literal["cold_mint", "mandatory_refresh", "advisory_refresh"]
class TokenExchangeMetricsSink(Protocol):
"""Observability seam for the exchange engine. Implementations must be best-effort: never raise
into the mint path, never block the calling thread, and never receive credential material --
``ExchangeError`` values are redacted by construction."""
def exchange_success(self, *, call_type: ExchangeCallType, duration_seconds: float) -> None: ...
def exchange_failure(
self, *, call_type: ExchangeCallType, duration_seconds: float, error: ExchangeError
) -> None: ...
def cache_hit(self) -> None: ...
class SyncTokenPoster(Protocol):
"""Returns the response for ANY status; never raises for status."""
def post(self, url: str, *, content: bytes, headers: Mapping[str, str], timeout: float) -> httpx.Response: ...

View file

@ -5,6 +5,7 @@ Utility functions for base LLM classes.
import copy
import json
from abc import ABC, abstractmethod
from collections.abc import Mapping
from typing import Any, Final
from openai.lib import _parsing, _pydantic
@ -65,6 +66,22 @@ class BaseLLMModelInfo(ABC):
"""
return []
def discover_models(
self, litellm_params: Mapping[str, object] | None = None
) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override
"""
Live model discovery for a configured deployment. Defaults to the api_key/api_base
facade every provider already implements via ``get_models``; a provider whose
discovery needs more of ``litellm_params`` (e.g. Anthropic's workload identity
federation) overrides this instead of widening ``get_models`` for every provider.
"""
api_key: Final = litellm_params.get("api_key") if litellm_params is not None else None
api_base: Final = litellm_params.get("api_base") if litellm_params is not None else None
return self.get_models(
api_key=api_key if isinstance(api_key, str) else None,
api_base=api_base if isinstance(api_base, str) else None,
)
@staticmethod
@abstractmethod
def get_api_key(api_key: str | None = None) -> str | None:

View file

@ -14,7 +14,7 @@ global state.
import re
from collections.abc import Mapping
from typing import Final
from typing import Final, Literal
from botocore.exceptions import (
CredentialRetrievalError,
@ -131,6 +131,10 @@ def is_mantle_claude_model(model: str) -> bool:
return "claude" in model.lower()
def mantle_health_check_mode(model: str) -> Literal["anthropic_messages"] | None:
return "anthropic_messages" if is_mantle_claude_model(model) else None
def mantle_supports_responses(model: str | None, model_cost: dict) -> bool:
"""Whether a Bedrock Mantle model can serve the native Responses API.

View file

@ -1382,10 +1382,12 @@ class HTTPHandler:
ssl_verify: bool | str | None = None,
disable_default_headers: bool
| None = False, # arize phoenix returns different API responses when user agent header in request
follow_redirects: bool = True,
):
self.timeout = timeout
self.ssl_verify = ssl_verify
self.disable_default_headers = disable_default_headers
self.follow_redirects = follow_redirects
self._owns_client = client is None
self._heal_lock = threading.Lock()
self._client = self.create_client() if client is None else client
@ -1410,7 +1412,7 @@ class HTTPHandler:
cert=cert,
headers=default_headers,
cookies=blocked_cookie_jar(),
follow_redirects=True,
follow_redirects=self.follow_redirects,
http2=http2_enabled(),
)
@ -1436,7 +1438,7 @@ class HTTPHandler:
self,
url: str,
params: dict | None = None,
headers: dict | None = None,
headers: Mapping[str, Any] | None = None,
follow_redirects: bool | None = None,
timeout: float | httpx.Timeout | None = None,
):

View file

@ -1,4 +1,5 @@
import asyncio
import inspect
import json
import ssl
from collections.abc import AsyncGenerator, AsyncIterator, Awaitable, Coroutine, Iterator, Mapping, Sequence
@ -18,6 +19,7 @@ from typing import (
Union,
cast,
get_type_hints,
runtime_checkable,
)
from urllib.parse import parse_qs, urlencode, urlparse, urlunparse
@ -277,6 +279,55 @@ class _MediaUploadKwargs(TypedDict, total=False):
timeout: float | httpx.Timeout
@runtime_checkable
class _AsyncFilesEnvironmentValidator(Protocol):
async def avalidate_environment(
self,
headers: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
model: str,
messages: list, # mutable-ok: mirrors the sync validate_environment contract this overrides
optional_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
litellm_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
api_key: str | None = None,
api_base: str | None = None,
) -> dict: ... # mutable-ok: mirrors the sync validate_environment contract this overrides
async def _avalidate_files_environment(
provider_config: BaseFilesConfig | BaseBatchesConfig,
*,
headers: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
model: str,
messages: list, # mutable-ok: mirrors the sync validate_environment contract this overrides
optional_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
litellm_params: dict, # mutable-ok: mirrors the sync validate_environment contract this overrides
api_key: str | None,
) -> dict: # mutable-ok: mirrors the sync validate_environment contract this overrides
"""Await the provider's async credential hook when it has one (e.g. Anthropic's workload
identity token exchange); otherwise offload the sync hook to a worker thread. Either way
the caller, an async file handler, never blocks the event loop on it."""
if isinstance(provider_config, _AsyncFilesEnvironmentValidator) and inspect.iscoroutinefunction(
provider_config.avalidate_environment
):
return await provider_config.avalidate_environment(
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
)
return await asyncio.to_thread(
provider_config.validate_environment,
headers=headers,
model=model,
messages=messages,
optional_params=optional_params,
litellm_params=litellm_params,
api_key=api_key,
)
class _SignedBodyKwargs(TypedDict, total=False):
data: ReadOnly[bytes]
json: ReadOnly[dict[str, object]]
@ -387,6 +438,21 @@ class _PreparedFileContentRequest(NamedTuple):
headers: dict
def _logged_file_content_request(
url: str,
params: dict,
request_headers: dict,
file_content_request: "FileContentRequest",
logging_obj: LiteLLMLoggingObj,
) -> _PreparedFileContentRequest:
logging_obj.pre_call(
input="",
api_key="",
additional_args={"api_base": url, "headers": request_headers, "file_id": file_content_request.get("file_id")},
)
return _PreparedFileContentRequest(url=url, params=params, headers=request_headers)
async def _aiter_bytes_then_close(response: httpx.Response, *, chunk_size: int) -> AsyncGenerator[bytes, None]:
try:
async for chunk in response.aiter_bytes(chunk_size=chunk_size):
@ -2027,7 +2093,7 @@ class BaseLLMHTTPHandler:
(
headers,
api_base,
) = anthropic_messages_provider_config.validate_anthropic_messages_environment(
) = await anthropic_messages_provider_config.avalidate_anthropic_messages_environment(
headers=merged_headers or {},
model=model,
messages=messages,
@ -3324,6 +3390,19 @@ class BaseLLMHTTPHandler:
"""
Creates a file using Gemini's two-step upload process
"""
if _is_async:
return self._avalidate_and_create_file(
create_file_data=create_file_data,
litellm_params=litellm_params,
provider_config=provider_config,
headers=headers,
api_base=api_base,
api_key=api_key,
logging_obj=logging_obj,
client=client,
timeout=timeout,
)
# get config from model, custom llm provider
headers = provider_config.validate_environment(
api_key=api_key,
@ -3353,18 +3432,6 @@ class BaseLLMHTTPHandler:
optional_params={},
)
if _is_async:
return self.async_create_file(
transformed_request=transformed_request,
litellm_params=litellm_params,
provider_config=provider_config,
headers=headers,
api_base=api_base,
logging_obj=logging_obj,
client=client,
timeout=timeout,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client()
else:
@ -3492,6 +3559,54 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params_with_url,
)
async def _avalidate_and_create_file(
self,
*,
create_file_data: CreateFileRequest,
litellm_params: dict, # mutable-ok: mirrors the create_file contract this dispatches for
provider_config: BaseFilesConfig,
headers: dict, # mutable-ok: mirrors the create_file contract this dispatches for
api_base: str | None,
api_key: str | None,
logging_obj: LiteLLMLoggingObj,
client: HTTPHandler | AsyncHTTPHandler | None,
timeout: float | httpx.Timeout | None,
) -> OpenAIFileObject:
validated_headers: Final = await _avalidate_files_environment(
provider_config,
headers=headers,
model="",
messages=[],
optional_params={},
litellm_params=litellm_params,
api_key=api_key,
)
complete_api_base: Final = provider_config.get_complete_file_url(
api_base=api_base,
api_key=api_key,
model="",
optional_params={},
litellm_params=litellm_params,
data=create_file_data,
)
if not complete_api_base:
raise ValueError("api_base is required for create_file")
return await self.async_create_file(
transformed_request=provider_config.transform_create_file_request(
model="",
create_file_data=create_file_data,
litellm_params=litellm_params,
optional_params={},
),
litellm_params=litellm_params,
provider_config=provider_config,
headers=validated_headers,
api_base=complete_api_base,
logging_obj=logging_obj,
client=client,
timeout=timeout,
)
async def async_create_file(
self,
transformed_request: Union[bytes, str, dict, "TwoStepFileUploadConfig"],
@ -3742,6 +3857,20 @@ class BaseLLMHTTPHandler:
if model is None:
raise ValueError("model is required for create_batch")
if _is_async:
return self._avalidate_and_create_batch(
create_batch_data=create_batch_data,
litellm_params=litellm_params,
provider_config=provider_config,
headers=headers,
api_base=api_base,
api_key=api_key,
logging_obj=logging_obj,
client=client,
timeout=timeout,
model=model,
)
headers = provider_config.validate_environment(
api_key=api_key,
headers=headers,
@ -3770,19 +3899,6 @@ class BaseLLMHTTPHandler:
optional_params={},
)
if _is_async:
return self.async_create_batch(
transformed_request=transformed_request,
litellm_params=litellm_params,
provider_config=provider_config,
headers=headers,
api_base=api_base,
logging_obj=logging_obj,
client=client,
timeout=timeout,
create_batch_data=create_batch_data,
)
if client is None or not isinstance(client, HTTPHandler):
sync_httpx_client = _get_httpx_client()
else:
@ -3920,6 +4036,56 @@ class BaseLLMHTTPHandler:
litellm_params=litellm_params,
)
async def _avalidate_and_create_batch(
self,
*,
create_batch_data: "CreateBatchRequest",
litellm_params: dict, # mutable-ok: mirrors the create_batch contract this dispatches for
provider_config: "BaseBatchesConfig",
headers: dict, # mutable-ok: mirrors the create_batch contract this dispatches for
api_base: str | None,
api_key: str | None,
logging_obj: "LiteLLMLoggingObj",
client: Union["HTTPHandler", "AsyncHTTPHandler"] | None,
timeout: float | httpx.Timeout | None,
model: str,
) -> "LiteLLMBatch":
validated_headers: Final = await _avalidate_files_environment(
provider_config,
headers=headers,
model=model,
messages=[],
optional_params={},
litellm_params=litellm_params,
api_key=api_key,
)
complete_api_base: Final = provider_config.get_complete_batch_url(
api_base=api_base,
api_key=api_key,
model=model,
optional_params={},
litellm_params=litellm_params,
data=create_batch_data,
)
if not complete_api_base:
raise ValueError("api_base is required for create_batch")
return await self.async_create_batch(
transformed_request=provider_config.transform_create_batch_request(
model=model,
create_batch_data=create_batch_data,
litellm_params=litellm_params,
optional_params={},
),
litellm_params=litellm_params,
provider_config=provider_config,
headers=validated_headers,
api_base=complete_api_base,
logging_obj=logging_obj,
client=client,
timeout=timeout,
create_batch_data=create_batch_data,
)
async def async_create_batch(
self,
transformed_request: bytes | str | dict,
@ -4521,7 +4687,8 @@ class BaseLLMHTTPHandler:
)
# Validate environment and get headers
headers = provider_config.validate_environment(
headers = await _avalidate_files_environment(
provider_config,
api_key=litellm_params.get("api_key"),
headers=headers,
model="",
@ -4646,7 +4813,8 @@ class BaseLLMHTTPHandler:
)
# Validate environment and get headers
headers = provider_config.validate_environment(
headers = await _avalidate_files_environment(
provider_config,
api_key=litellm_params.get("api_key"),
headers=headers,
model="",
@ -4770,7 +4938,8 @@ class BaseLLMHTTPHandler:
)
# Validate environment and get headers
headers = provider_config.validate_environment(
headers = await _avalidate_files_environment(
provider_config,
api_key=litellm_params.get("api_key"),
headers=headers,
model="",
@ -4958,7 +5127,7 @@ class BaseLLMHTTPHandler:
else:
async_httpx_client = client
prepared: Final = self._prepare_file_content_request(
prepared: Final = await self._aprepare_file_content_request(
file_content_request=file_content_request,
provider_config=provider_config,
litellm_params=litellm_params,
@ -5004,7 +5173,7 @@ class BaseLLMHTTPHandler:
client if client is not None else get_async_httpx_client(llm_provider=provider_config.custom_llm_provider)
)
prepared: Final = self._prepare_file_content_request(
prepared: Final = await self._aprepare_file_content_request(
file_content_request=file_content_request,
provider_config=provider_config,
litellm_params=litellm_params,
@ -5062,16 +5231,31 @@ class BaseLLMHTTPHandler:
optional_params={},
litellm_params=litellm_params,
)
logging_obj.pre_call(
input="",
api_key="",
additional_args={
"api_base": url,
"headers": request_headers,
"file_id": file_content_request.get("file_id"),
},
return _logged_file_content_request(url, params, request_headers, file_content_request, logging_obj)
@staticmethod
async def _aprepare_file_content_request(
file_content_request: "FileContentRequest",
provider_config: BaseFilesConfig,
litellm_params: dict,
headers: dict,
logging_obj: LiteLLMLoggingObj,
) -> "_PreparedFileContentRequest":
url, params = provider_config.transform_file_content_request(
file_content_request=file_content_request,
optional_params={},
litellm_params=litellm_params,
)
return _PreparedFileContentRequest(url=url, params=params, headers=request_headers)
request_headers: Final = await _avalidate_files_environment(
provider_config,
api_key=litellm_params.get("api_key"),
headers=headers,
model="",
messages=[],
optional_params={},
litellm_params=litellm_params,
)
return _logged_file_content_request(url, params, request_headers, file_content_request, logging_obj)
def _prepare_fake_stream_request(
self,

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