ci: merge main into litellm_chore_f0413c

This commit is contained in:
yuneng 2026-10-04 00:57:44 +00:00
commit 822a9ddbc7
246 changed files with 24814 additions and 2519 deletions

View file

@ -15,6 +15,9 @@ on:
- gateway/main.py
- backend/Dockerfile
- backend/main.py
- deploy/lens/**
- litellm/proxy/lens/release.py
- tests/e2e/migrations/lens_compose_smoke.sh
- docker/component_entrypoint.sh
- docker/entrypoint.sh
- litellm/proxy/prisma_migration.py
@ -113,7 +116,7 @@ jobs:
persist-credentials: false
- name: Build runtime image
run: docker build -f Dockerfile -t litellm-runtime-scan:${{ github.sha }} .
run: docker build --build-arg LITELLM_RELEASE_TAG=v0.0.0-lens-ci -f Dockerfile -t litellm-runtime-scan:${{ github.sha }} .
- name: Set up Python
uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0
@ -127,6 +130,11 @@ jobs:
python -m pip install "pytest==9.0.3"
python -m pytest tests/proxy_migration_tests/test_offline_image_migration.py tests/proxy_migration_tests/test_image_bedrock_realtime_extra.py -v
- name: Verify the bundled Lens Compose installation and restart
env:
LITELLM_IMAGE: litellm-runtime-scan:${{ github.sha }}
run: bash tests/e2e/migrations/lens_compose_smoke.sh
migrations-image:
name: migrations-image
runs-on: ubuntu-latest

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 \

View file

@ -326,6 +326,7 @@ jobs:
tests/unit/enterprise/proxy/hooks
tests/unit/enterprise/proxy/management_endpoints
tests/unit/enterprise/proxy/test_audit_logging_endpoints.py
tests/unit/enterprise/proxy/test_liteadmin.py
tests/unit/enterprise/enterprise_callbacks/test_prometheus_logging_callbacks.py
workers: 4
reruns: 2

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

@ -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,10 @@
FROM python:3.12-slim
ARG LITELLM_RELEASE_TAG=""
RUN : "${LITELLM_RELEASE_TAG:?Pass --build-arg LITELLM_RELEASE_TAG matching the gateway}"
ENV LITELLM_RELEASE_TAG=${LITELLM_RELEASE_TAG}
WORKDIR /app
RUN pip install --no-cache-dir httpx==0.28.1 pydantic==2.11.7
COPY litellm/proxy/lens/__init__.py litellm/proxy/lens/models.py litellm/proxy/lens/trace_store.py litellm/proxy/lens/analysis.py litellm/proxy/lens/worker.py /app/lens/
COPY litellm/proxy/lens/__init__.py litellm/proxy/lens/models.py litellm/proxy/lens/trace_store.py litellm/proxy/lens/analysis.py litellm/proxy/lens/worker.py litellm/proxy/lens/release.py /app/lens/
COPY litellm/proxy/lens/prompts/ /app/lens/prompts/
USER 65532:65532
CMD ["python", "-m", "lens.worker"]

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
```
The worker image defaults to the chart's application version, and the chart connects it to the backend service. Keep these values and the Secret when upgrading the chart so the gateway and worker upgrade together. `lensWorker.replicaCount` controls simultaneous investigations. To use a private registry or external proxy, set `lensWorker.image.repository`, `lensWorker.image.tag`, and `lensWorker.url`. The dashboard uses the chart's worker image for standalone install commands too
## Standalone worker
Upgrade your existing LiteLLM proxy to a release that includes Lens with PostgreSQL and agent tracing. Configure one ClickHouse URL for trace writes, bounded reads, and Lens queries:
@ -23,17 +76,17 @@ In **Lens > Investigations**, click **Connect worker**, choose an analysis model
The command already contains the compatible worker image and one worker token. The selected virtual key stays on the proxy; its secret is never sent to the worker. No source checkout, environment file, or second LiteLLM deployment is needed. Keep the command private because it includes the token. The LiteLLM release provides the dashboard and APIs; the container only runs background analysis
The dashboard and Compose file pin a verified worker image by digest. The image uses Linux amd64, and the generated command selects that platform. CI also publishes immutable `:sha-<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
@ -130,3 +183,14 @@ The Lens API now uses `/lens` instead of `/engine`, list responses use `lenses`,
Stop workers and let active scans finish before upgrading. Deploy proxy instances together: older proxies cannot use the renamed database tables. The schema migration renames the three Lens tables and the run-history identifier column in place, preserving saved investigations, findings, history, worker credentials, and billing assignments. Existing migration files retain their original names and checksums
Upgrades using `--use_prisma_db_push` stop before schema changes if any legacy Lens table exists, preventing Prisma from dropping saved data. Apply `litellm-proxy-extras/litellm_proxy_extras/migrations/20261001100000_rename_lens/migration.sql` to the configured database schema before retrying. Deployments already using migration history can instead start without `--use_prisma_db_push` to apply the shipped migration normally. Fresh databases and databases already using the renamed tables can continue using database push
## Release compatibility
Released gateway and worker images carry `LITELLM_RELEASE_TAG`. A worker announces its release and protocol before claiming an investigation. A mismatch returns HTTP 409 with the required image, leaving queued investigations untouched. During a rolling upgrade, workers wait for a gateway from their release
The dashboard reads its image from the running gateway. `LENS_WORKER_IMAGE` overrides the registry/image for private deployments. Worker-only Compose accepts `LITELLM_VERSION` (without `v`) or an explicit `LENS_WORKER_IMAGE`. Release workers are available as `ghcr.io/berriai/litellm-lens-worker:vX.Y.Z` and `docker.io/litellm/litellm-lens-worker:vX.Y.Z`, including matching RC/dev suffixes, on amd64 and arm64
For source development, use `make lens-dev`, which gives the proxy and source worker the same commit identity. For custom containers, build both from the same checkout with `--build-arg LITELLM_RELEASE_TAG=sha-$(git rev-parse HEAD)` and set the proxy's `LENS_WORKER_IMAGE` to the worker image you built. An unlabelled custom build refuses worker setup and claims instead of guessing from the Python package version. Normal package-index installations use their installed release version
The hourly development pipeline pins all component images to the same selected commit and publishes its chart only after every build and worker smoke test succeeds. The public commit-tagged worker workflow publishes on Lens-related changes, so an arbitrary `main` commit may require building your own pair; do not substitute the newest available worker

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

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,13 @@ Directory of the collector's unix socket, shared by the gateway and
collector containers through an emptyDir. Empty when the sidecar is off
or gateway.collector.address is a tcp://127.0.0.1:<port> address.
*/}}
{{- define "litellm.lensWorker.image" -}}
{{- $backendTag := .Values.backend.image.tag | default .Chart.AppVersion -}}
{{- $releaseTag := ternary (printf "v%s" $backendTag) $backendTag (regexMatch "^[0-9]" $backendTag) -}}
{{- $tag := .Values.lensWorker.image.tag | default $releaseTag -}}
{{- printf "%s:%s" .Values.lensWorker.image.repository $tag -}}
{{- end -}}
{{- define "litellm.gateway.collectorSocketDir" -}}
{{- if and .Values.gateway.collector.enabled (hasPrefix "unix://" .Values.gateway.collector.address) -}}
{{- dir (trimPrefix "unix://" .Values.gateway.collector.address) -}}

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,114 @@
suite: Lens worker release and credentials
templates:
- lens/deployment.yaml
- backend/deployment.yaml
- gateway/configmap.yaml
values:
- ./values/required.yaml
tests:
- it: keeps the worker opt in
template: lens/deployment.yaml
asserts:
- hasDocuments:
count: 0
- it: requires a limited worker credential when enabled
template: lens/deployment.yaml
set:
lensWorker.enabled: true
asserts:
- failedTemplate:
errorMessage: lensWorker.tokenSecret.name must reference a Lens worker token
- it: uses the chart release and a secret without granting Kubernetes access
template: lens/deployment.yaml
chart:
appVersion: v1.2.3
set:
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker:v1.2.3
- equal:
path: spec.template.spec.containers[0].env[1].valueFrom.secretKeyRef
value:
name: lens-credential
key: token
- equal:
path: spec.template.spec.automountServiceAccountToken
value: false
- equal:
path: spec.template.spec.containers[0].securityContext.readOnlyRootFilesystem
value: true
- equal:
path: spec.template.spec.volumes[0].emptyDir
value:
medium: Memory
sizeLimit: 1Gi
- it: advertises the same private dev image to standalone installers
template: backend/deployment.yaml
set:
lensWorker.image.repository: registry.example/lens-worker
lensWorker.image.tag: branch-main-1234567
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LENS_WORKER_IMAGE
value: registry.example/lens-worker:branch-main-1234567
- it: supports an external gateway and a registry override
template: lens/deployment.yaml
set:
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
lensWorker.url: https://gateway.example/proxy
lensWorker.image.repository: registry.example/lens-worker
lensWorker.image.tag: branch-main-1234567
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: registry.example/lens-worker:branch-main-1234567
- equal:
path: spec.template.spec.containers[0].env[0].value
value: https://gateway.example/proxy
- it: prefixes a numeric chart release with v
template: lens/deployment.yaml
chart:
appVersion: 1.2.3-rc.4
set:
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-rc.4
- it: follows a backend image override when no worker tag is set
template: lens/deployment.yaml
set:
backend.image.tag: branch-main-1234567
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker:branch-main-1234567
- it: recommends the overridden backend release for standalone installers
template: backend/deployment.yaml
set:
backend.image.tag: v1.2.3-dev.4
asserts:
- contains:
path: spec.template.spec.containers[0].env
content:
name: LENS_WORKER_IMAGE
value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-dev.4
- it: normalizes a numeric backend tag to the published worker tag
template: lens/deployment.yaml
set:
backend.image.tag: 1.2.3-dev.4
lensWorker.enabled: true
lensWorker.tokenSecret.name: lens-credential
asserts:
- equal:
path: spec.template.spec.containers[0].image
value: ghcr.io/berriai/litellm-lens-worker:v1.2.3-dev.4

View file

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

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

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

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

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

@ -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 {
@ -359,7 +359,11 @@ fn unrelated_sibling_transport_leaves_cost_unchanged(
let calls: Vec<_> = rows
.iter()
.filter(|row| {
row.kind == ObservationType::Llm && !row.call_keys.contains(&CallKey::Transport)
row.kind == ObservationType::Llm
&& !row
.call_keys
.iter()
.any(|key| matches!(key, CallKey::Transport | CallKey::GatewayAttempt))
})
.cloned()
.collect();
@ -408,7 +412,9 @@ fn redundant_genai_response_id_keeps_call_evidence(
.iter()
.filter_map(|key| match key {
CallKey::ProviderResponse(id) => Some(id.clone()),
CallKey::LiteLlmRequest(_) | CallKey::Transport => None,
CallKey::LiteLlmRequest(_) | CallKey::Transport | CallKey::GatewayAttempt => {
None
}
})
.collect();
(!response_ids.is_empty()).then(|| {

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

@ -255,6 +255,7 @@ fn llamaindex_wrapped_responses_keep_provider_call_keys(#[case] body: &[u8]) {
#[case::request(litellm_traces::CallKey::LiteLlmRequest("request:with:colons".to_owned()))]
#[case::response(litellm_traces::CallKey::ProviderResponse("response:with:colons".to_owned()))]
#[case::transport(litellm_traces::CallKey::Transport)]
#[case::gateway_attempt(litellm_traces::CallKey::GatewayAttempt)]
fn call_keys_round_trip_through_storage(#[case] key: litellm_traces::CallKey) {
assert_eq!(
key.to_string().parse::<litellm_traces::CallKey>().unwrap(),
@ -272,6 +273,8 @@ fn call_keys_round_trip_through_storage(#[case] key: litellm_traces::CallKey) {
#[case::missing_response("provider_response:")]
#[case::missing_request("litellm_request:")]
#[case::transport_id("transport:unexpected")]
#[case::gateway_attempt_separator("gateway_attempt")]
#[case::gateway_attempt_id("gateway_attempt:unexpected")]
#[case::unknown("unknown:id")]
fn malformed_call_keys_are_rejected_at_the_boundary(#[case] encoded: &str) {
assert!(encoded.parse::<litellm_traces::CallKey>().is_err());

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,

View file

@ -58,6 +58,7 @@ from litellm.types.utils import (
from litellm.utils import convert_to_model_response_object
from ..common_utils import OpenAIError
from ..workload_identity import get_workload_identity_bearer_token, resolve_openai_workload_identity_config
if TYPE_CHECKING:
from litellm.litellm_core_utils.litellm_logging import Logging as _LiteLLMLoggingObj
@ -73,6 +74,11 @@ else:
_NO_TOOLS_UPDATE: Final[Mapping[str, object]] = MappingProxyType({})
def _litellm_params_str(litellm_params: Mapping[str, object] | None, key: str) -> str | None:
value: Final = litellm_params.get(key) if litellm_params is not None else None
return value if isinstance(value, str) else None
class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
"""
Reference: https://platform.openai.com/docs/api-reference/chat/create
@ -765,28 +771,39 @@ class OpenAIGPTConfig(BaseLLMModelInfo, BaseConfig):
"""
Calls OpenAI's `/v1/models` endpoint and returns the list of models.
"""
if api_base is None:
api_base = "https://api.openai.com"
if api_key is None:
api_key = get_secret_str("OPENAI_API_KEY")
# Strip api_base to just the base URL (scheme + host + port)
parsed_url: Final = httpx.URL(api_base)
base_url = f"{parsed_url.scheme}://{parsed_url.host}"
if parsed_url.port:
base_url += f":{parsed_url.port}"
response: Final = litellm.module_level_client.get(
url=f"{base_url}/v1/models",
headers={"Authorization": f"Bearer {api_key}"},
return self._fetch_model_ids(
api_base=api_base, bearer_token=get_secret_str("OPENAI_API_KEY") if api_key is None else api_key
)
def discover_models(
self, litellm_params: Mapping[str, object] | None = None
) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override
if type(self) is not OpenAIGPTConfig:
return super().discover_models(litellm_params)
api_key: Final = _litellm_params_str(litellm_params, "api_key")
api_base: Final = _litellm_params_str(litellm_params, "api_base")
workload_identity_config: Final = resolve_openai_workload_identity_config(
api_key=api_key, api_base=api_base, litellm_params=litellm_params
)
if workload_identity_config is None:
return self.get_models(api_key=api_key, api_base=api_base)
return self._fetch_model_ids(
api_base=api_base, bearer_token=get_workload_identity_bearer_token(workload_identity_config)
)
@staticmethod
def _fetch_model_ids(
api_base: str | None, bearer_token: str | None
) -> list[str]: # mutable-ok: matches get_models' list[str] contract shared by every provider override
parsed_url: Final = httpx.URL(api_base or "https://api.openai.com")
port_suffix: Final = f":{parsed_url.port}" if parsed_url.port else ""
response: Final = litellm.module_level_client.get(
url=f"{parsed_url.scheme}://{parsed_url.host}{port_suffix}/v1/models",
headers={"Authorization": f"Bearer {bearer_token}"},
)
if response.status_code != 200:
raise Exception(f"Failed to get models: {response.text}")
models: Final = response.json()["data"]
return [model["id"] for model in models]
return [model["id"] for model in response.json()["data"]]
@staticmethod
def get_api_key(api_key: str | None = None) -> str | None:

View file

@ -382,8 +382,11 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
organization: str | None = None,
client: OpenAI | AsyncOpenAI | None = None,
shared_session: Optional["ClientSession"] = None,
litellm_params: Mapping[str, object] | None = None,
) -> OpenAI | AsyncOpenAI | None:
workload_identity_config: Final = resolve_openai_workload_identity_config(api_key=api_key, api_base=api_base)
workload_identity_config: Final = resolve_openai_workload_identity_config(
api_key=api_key, api_base=api_base, litellm_params=litellm_params
)
client_initialization_params: Final[dict] = locals()
if client is None:
if not isinstance(max_retries, int):
@ -773,6 +776,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=max_retries,
organization=organization,
stream_options=stream_options,
litellm_params=litellm_params,
)
else:
if not isinstance(max_retries, int):
@ -786,6 +790,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=max_retries,
organization=organization,
client=client,
litellm_params=litellm_params,
)
## LOGGING
@ -928,6 +933,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
organization=organization,
client=client,
shared_session=shared_session,
litellm_params=litellm_params,
)
## LOGGING
@ -1024,6 +1030,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=None,
headers=None,
stream_options: dict | None = None,
litellm_params: Mapping[str, object] | None = None,
):
data["stream"] = True
data.update(self.get_stream_options(stream_options=stream_options, api_base=api_base))
@ -1037,6 +1044,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=max_retries,
organization=organization,
client=client,
litellm_params=litellm_params,
)
## LOGGING
logging_obj.pre_call(
@ -1109,6 +1117,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
organization=organization,
client=client,
shared_session=shared_session,
litellm_params=litellm_params,
)
## LOGGING
logging_obj.pre_call(
@ -1243,6 +1252,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client: AsyncOpenAI | None = None,
max_retries=None,
shared_session: Optional["ClientSession"] = None,
litellm_params: Mapping[str, object] | None = None,
):
try:
openai_aclient: Final[AsyncOpenAI] = self._get_openai_client(
@ -1253,6 +1263,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
max_retries=max_retries,
client=client,
shared_session=shared_session,
litellm_params=litellm_params,
)
raw_response: Final = await self.make_openai_embedding_request(
openai_aclient=openai_aclient,
@ -1316,6 +1327,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
aembedding=None,
max_retries: int | None = None,
shared_session: Optional["ClientSession"] = None,
litellm_params: Mapping[str, object] | None = None,
) -> EmbeddingResponse:
super().embedding()
try:
@ -1342,6 +1354,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
client=client,
max_retries=max_retries,
shared_session=shared_session,
litellm_params=litellm_params,
)
openai_client: Final[OpenAI] = self._get_openai_client(
@ -1351,6 +1364,7 @@ class OpenAIChatCompletion(BaseLLM, BaseOpenAILLM):
timeout=timeout,
max_retries=max_retries,
client=client,
litellm_params=litellm_params,
)
## embedding CALL

View file

@ -30,6 +30,7 @@ from litellm.types.llms.openai import *
from litellm.types.responses.main import *
from litellm.types.router import GenericLiteLLMParams
from litellm.types.utils import LlmProviders
from litellm.types.workload_identity import OPENAI_WIF_KWARGS_KEYS
from ..common_utils import OpenAIError
from ..workload_identity import get_workload_identity_bearer_token, resolve_openai_workload_identity_config
@ -599,7 +600,11 @@ class OpenAIResponsesAPIConfig(BaseResponsesAPIConfig):
api_key = litellm_params.api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY")
headers.setdefault("Content-Type", "application/json")
workload_identity_config: Final = (
resolve_openai_workload_identity_config(api_key=api_key, api_base=litellm_params.api_base)
resolve_openai_workload_identity_config(
api_key=api_key,
api_base=litellm_params.api_base,
litellm_params=litellm_params.model_dump(include=set(OPENAI_WIF_KWARGS_KEYS)),
)
if self.custom_llm_provider is LlmProviders.OPENAI
else None
)

View file

@ -1,5 +1,6 @@
from __future__ import annotations
from collections.abc import Mapping
from dataclasses import dataclass
from functools import lru_cache
from typing import TYPE_CHECKING, Final
@ -43,6 +44,7 @@ class OpenAIWorkloadIdentityConfig:
def resolve_openai_workload_identity_config(
api_key: str | None,
api_base: str | None,
litellm_params: Mapping[str, object] | None = None,
) -> OpenAIWorkloadIdentityConfig | None:
static_api_key: Final = normalize_nonempty_secret_str(api_key) or normalize_nonempty_secret_str(
get_secret_str("OPENAI_API_KEY")
@ -54,10 +56,12 @@ def resolve_openai_workload_identity_config(
)
if not _targets_openai_api(effective_api_base):
return None
identity_provider_id: Final = get_secret_str("OPENAI_IDENTITY_PROVIDER_ID")
service_account_id: Final = get_secret_str("OPENAI_SERVICE_ACCOUNT_ID")
token_file: Final = get_secret_str("OPENAI_IDENTITY_TOKEN_FILE")
if not identity_provider_id or not service_account_id or not token_file:
identity_provider_id: Final = _config_value(
litellm_params, "openai_identity_provider_id", "OPENAI_IDENTITY_PROVIDER_ID"
)
service_account_id: Final = _config_value(litellm_params, "openai_service_account_id", "OPENAI_SERVICE_ACCOUNT_ID")
token_file: Final = _config_value(litellm_params, "openai_identity_token_file", "OPENAI_IDENTITY_TOKEN_FILE")
if identity_provider_id is None or service_account_id is None or token_file is None:
return None
return OpenAIWorkloadIdentityConfig(
identity_provider_id=identity_provider_id,
@ -77,6 +81,13 @@ async def get_workload_identity_bearer_token_for_api_base(api_base: str) -> str
return await _workload_identity_auth(config).get_token_async()
def _config_value(litellm_params: Mapping[str, object] | None, param_key: str, env_name: str) -> str | None:
param_value: Final = litellm_params.get(param_key) if litellm_params is not None else None
if isinstance(param_value, str) and param_value:
return param_value
return normalize_nonempty_secret_str(get_secret_str(env_name))
def _targets_openai_api(api_base: str | None) -> bool:
if api_base is None:
return True

View file

@ -1240,6 +1240,9 @@ class VertexAITokenCounter(BaseTokenCounter):
) -> TokenCountResponse | None:
import copy
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
VertexAIError as PartnerVertexAIError,
)
from litellm.llms.vertex_ai.vertex_ai_partner_models.main import (
VertexAIPartnerModels,
)
@ -1269,14 +1272,32 @@ class VertexAITokenCounter(BaseTokenCounter):
"vertex_ai_credentials"
)
result = await partner_models_handler.count_tokens(
model=model_to_use,
messages=messages or [],
litellm_params=partner_litellm_params,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_credentials=vertex_credentials,
)
try:
result = await partner_models_handler.count_tokens(
model=model_to_use,
messages=messages or [],
litellm_params=partner_litellm_params,
vertex_project=vertex_project,
vertex_location=vertex_location,
vertex_credentials=vertex_credentials,
system=system,
tools=tools,
)
except (PartnerVertexAIError, httpx.HTTPStatusError) as e:
status_code: Final = e.response.status_code
error_message: Final = e.message if isinstance(e, PartnerVertexAIError) else e.response.text
verbose_logger.warning(
"Vertex AI partner CountTokens API error: status=%s, message=%s", status_code, error_message
)
return TokenCountResponse(
total_tokens=0,
request_model=request_model,
model_used=model_to_use,
tokenizer_type="vertex_ai_partner_models",
error=True,
error_message=error_message,
status_code=status_code,
)
if result is not None:
return TokenCountResponse(

View file

@ -2,6 +2,7 @@
## API Handler for calling Vertex AI Partner Models
from collections.abc import Callable
from enum import Enum
from types import MappingProxyType
from typing import Final
import httpx
@ -263,6 +264,8 @@ class VertexAIPartnerModels(VertexBase):
vertex_project=None,
vertex_location=None,
vertex_credentials=None,
system: object | None = None,
tools: list[dict[str, object]] | None = None,
):
"""
Count tokens for Vertex AI partner models (Anthropic Claude, Mistral, etc.)
@ -296,6 +299,9 @@ class VertexAIPartnerModels(VertexBase):
request_data: Final = {
"model": model,
"messages": messages,
**MappingProxyType(
{key: value for key, value in (("system", system), ("tools", tools)) if value is not None}
),
}
# Prepare litellm_params with credentials

View file

@ -144,6 +144,7 @@ from litellm.types.utils import (
RawRequestTypedDict,
StreamingChoices,
)
from litellm.types.workload_identity import ANTHROPIC_WIF_KWARGS_KEYS, OPENAI_WIF_KWARGS_KEYS
from litellm.utils import (
Choices,
CustomStreamWrapper,
@ -5489,7 +5490,9 @@ def completion(
api_base=api_base,
api_key=api_key,
litellm_params=(
GenericLiteLLMParams(**_supplemental_provider_params) if _supplemental_provider_params else None
GenericLiteLLMParams.model_validate(_supplemental_provider_params)
if _supplemental_provider_params
else None
),
)
@ -5716,7 +5719,12 @@ def completion(
gigachat_access_token=kwargs.get("gigachat_access_token"),
**{
key: kwargs[key]
for key in (*AWS_CREDENTIAL_KWARGS_KEYS, PROVIDER_AFFINITY_HEADER_KWARG_KEY)
for key in (
*AWS_CREDENTIAL_KWARGS_KEYS,
*ANTHROPIC_WIF_KWARGS_KEYS,
*OPENAI_WIF_KWARGS_KEYS,
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
)
if key in kwargs
},
)
@ -6553,6 +6561,7 @@ def embedding(
aembedding=aembedding,
max_retries=max_retries,
shared_session=shared_session,
litellm_params=litellm_params_dict,
)
elif custom_llm_provider == "databricks":
api_base = api_base or litellm.api_base or get_secret("DATABRICKS_API_BASE")
@ -7802,7 +7811,7 @@ async def amoderation(
# only supports open ai for now
api_key = api_key or litellm.api_key or litellm.openai_key or get_secret_str("OPENAI_API_KEY")
optional_params: Final = GenericLiteLLMParams(**kwargs)
optional_params: Final = GenericLiteLLMParams.model_validate(kwargs)
litellm_logging_obj: Final[LiteLLMLoggingObj | None] = kwargs.get("litellm_logging_obj", None)
_dynamic_api_base = None
try:
@ -8519,7 +8528,7 @@ def speech(
VertexAITextToSpeechConfig,
)
generic_optional_params: Final = GenericLiteLLMParams(**kwargs)
generic_optional_params: Final = GenericLiteLLMParams.model_validate(kwargs)
# Handle Gemini models separately (they use speech_to_completion_bridge)
if "gemini" in model:
@ -8725,7 +8734,7 @@ def speech(
async def ahealth_check(
model_params: dict,
mode: str | None = "chat",
mode: str | None = None,
prompt: str | None = None,
input: list | None = None,
):
@ -8740,7 +8749,8 @@ async def ahealth_check(
}
"""
from litellm.litellm_core_utils.cached_imports import get_litellm_logging_class
from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers
from litellm.litellm_core_utils.health_check_helpers import HealthCheckHelpers, default_health_check_mode
from litellm.litellm_core_utils.health_check_utils import OPTIONAL_STR
# Use cached import helper to lazy-load Logging class (only loads when function is called)
Logging: Final = get_litellm_logging_class()
@ -8765,28 +8775,25 @@ async def ahealth_check(
)
#########################################################
try:
model: str | None = model_params.get("model", None)
if model is None:
requested_model: Final = OPTIONAL_STR.validate_python(model_params.get("model", None))
if requested_model is None:
raise Exception("model not set")
if model in litellm.model_cost and mode is None:
mode = litellm.model_cost[model].get("mode")
custom_llm_provider_from_params: Final = model_params.get("custom_llm_provider", None)
api_base_from_params: Final = model_params.get("api_base", None)
api_key_from_params: Final = model_params.get("api_key", None)
model, custom_llm_provider, _, _ = get_llm_provider(
model=model,
model=requested_model,
custom_llm_provider=custom_llm_provider_from_params,
api_base=api_base_from_params,
api_key=api_key_from_params,
)
if model in litellm.model_cost and mode is None:
mode = litellm.model_cost[model].get("mode")
model_params["cache"] = {"no-cache": True} # don't used cached responses for making health check calls
mode = mode or "chat"
mode = mode or default_health_check_mode(
requested_model=requested_model, model=model, custom_llm_provider=custom_llm_provider
)
if "*" in model:
return await HealthCheckHelpers.ahealth_check_wildcard_models(
model=model,
@ -8815,12 +8822,6 @@ async def ahealth_check(
if isinstance(stack_trace, str):
stack_trace = stack_trace[:1000]
if mode is None:
return {
"error": f"error:{e}. Missing `mode`. Set the `mode` for the model - https://docs.litellm.ai/docs/proxy/health#embedding-models \nstacktrace: {stack_trace}",
"exception": e,
}
error_to_return: Final = str(e) + "\nstack trace: " + stack_trace
raw_request_typed_dict: Final = litellm_logging_obj.model_call_details.get("raw_request_typed_dict")

View file

@ -10841,31 +10841,31 @@
"deprecation_date": "2028-02-09",
"input_cost_per_token": 1.3e-07,
"litellm_provider": "azure",
"max_input_tokens": 8191,
"max_tokens": 8191,
"max_input_tokens": 8192,
"max_tokens": 8192,
"mode": "embedding",
"output_cost_per_token": 0.0,
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure"
},
"azure/text-embedding-3-small": {
"deprecation_date": "2028-02-09",
"input_cost_per_token": 2e-08,
"litellm_provider": "azure",
"max_input_tokens": 8191,
"max_tokens": 8191,
"max_input_tokens": 8192,
"max_tokens": 8192,
"mode": "embedding",
"output_cost_per_token": 0.0,
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure"
},
"azure/text-embedding-ada-002": {
"deprecation_date": "2028-02-09",
"input_cost_per_token": 1e-07,
"litellm_provider": "azure",
"max_input_tokens": 8191,
"max_tokens": 8191,
"max_input_tokens": 8192,
"max_tokens": 8192,
"mode": "embedding",
"output_cost_per_token": 0.0,
"source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'"
"source": "https://learn.microsoft.com/en-us/azure/foundry/foundry-models/concepts/models-sold-directly-by-azure"
},
"azure/speech/azure-tts": {
"input_cost_per_character": 1.5e-05,
@ -42177,14 +42177,13 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-chat-v3-0324": {
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 2.9e-07,
"input_cost_per_token": 2.5e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 163840,
"max_output_tokens": 147456,
"max_tokens": 147456,
"mode": "chat",
"output_cost_per_token": 1.14e-06,
"output_cost_per_token": 1e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -42349,15 +42348,14 @@
"supports_web_search": false
},
"openrouter/deepseek/deepseek-v4-pro-0813": {
"cache_read_input_token_cost": 4.4e-08,
"input_cost_per_token": 1.32e-06,
"cache_read_input_token_cost": 1.4e-07,
"input_cost_per_token": 2.2e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 393216,
"max_tokens": 393216,
"max_output_tokens": 943718,
"max_tokens": 943718,
"mode": "chat",
"off_peak_pricing": {"input_cost_per_token":6.6e-7,"output_cost_per_token":0.00000198,"cache_read_input_token_cost":2.2e-8,"windows":[{"hours_utc":"00:00-00:00","weekdays":["saturday","sunday"]},{"hours_utc":"00:00-01:00","weekdays":["monday","tuesday","wednesday","thursday","friday"]},{"hours_utc":"04:00-06:00","weekdays":["monday","tuesday","wednesday","thursday","friday"]},{"hours_utc":"10:00-00:00","weekdays":["monday","tuesday","wednesday","thursday","friday"]}]},
"output_cost_per_token": 3.96e-06,
"output_cost_per_token": 4.2e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -67734,8 +67732,8 @@
"supports_prompt_caching": true
},
"openrouter/deepseek/deepseek-v4-flash-0731": {
"cache_read_input_token_cost": 5.1e-09,
"input_cost_per_token": 5.1e-09,
"cache_read_input_token_cost": 1.37e-08,
"input_cost_per_token": 1.52e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 943718,
@ -67825,14 +67823,14 @@
"supports_web_search": false
},
"openrouter/moonshotai/kimi-k3": {
"cache_read_input_token_cost": 2.7e-07,
"input_cost_per_token": 2.7e-06,
"cache_read_input_token_cost": 4.9e-07,
"input_cost_per_token": 4.99e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 943718,
"max_tokens": 943718,
"mode": "chat",
"output_cost_per_token": 1.35e-05,
"output_cost_per_token": 1.3e-05,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -67951,13 +67949,13 @@
},
"openrouter/z-ai/glm-5.2": {
"cache_read_input_token_cost": 2.6e-07,
"input_cost_per_token": 4.1e-07,
"input_cost_per_token": 3e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 1048576,
"max_output_tokens": 943718,
"max_tokens": 943718,
"mode": "chat",
"output_cost_per_token": 3.99e-06,
"output_cost_per_token": 3.49e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -68333,24 +68331,24 @@
"supports_web_search": false
},
"openrouter/moonshotai/kimi-k2.6": {
"input_cost_per_token": 4.3415e-07,
"output_cost_per_token": 1.828e-06,
"cache_read_input_token_cost": 7.312e-08,
"cache_read_input_token_cost": 1.6e-07,
"input_cost_per_token": 9.5e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 235929,
"max_tokens": 235929,
"mode": "chat",
"output_cost_per_token": 4e-06,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_tool_choice": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_parallel_function_calling": true,
"supports_pdf_input": false,
"supports_vision": true,
"supports_prompt_caching": true,
"supports_reasoning": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": true,
"supports_web_search": false
},
"openrouter/google/gemma-4-26b-a4b-it": {
@ -69370,13 +69368,13 @@
"supports_web_search": false
},
"openrouter/qwen/qwen3-30b-a3b-instruct-2507": {
"input_cost_per_token": 1e-07,
"input_cost_per_token": 4.815e-08,
"litellm_provider": "openrouter",
"max_input_tokens": 262144,
"max_output_tokens": 235929,
"max_tokens": 235929,
"max_output_tokens": 32000,
"max_tokens": 32000,
"mode": "chat",
"output_cost_per_token": 3e-07,
"output_cost_per_token": 1.9305e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
@ -69933,21 +69931,22 @@
"supports_web_search": false
},
"openrouter/meta-llama/llama-3.3-70b-instruct": {
"input_cost_per_token": 1e-07,
"output_cost_per_token": 3.2e-07,
"cache_read_input_token_cost": 1.1e-07,
"input_cost_per_token": 2.2e-07,
"litellm_provider": "openrouter",
"max_input_tokens": 131072,
"max_output_tokens": 16384,
"max_tokens": 16384,
"mode": "chat",
"output_cost_per_token": 5e-07,
"source": "https://openrouter.ai/api/v1/models",
"supports_audio_input": false,
"supports_function_calling": true,
"supports_pdf_input": false,
"supports_prompt_caching": false,
"supports_reasoning": false,
"supports_tool_choice": true,
"supports_response_schema": true,
"supports_tool_choice": true,
"supports_vision": false,
"supports_web_search": false
},
@ -79983,6 +79982,7 @@
"supports_web_search": false
},
"azure_ai/kimi-k2-thinking": {
"deprecation_date": "2026-03-29",
"input_cost_per_token": 6e-07,
"litellm_provider": "azure_ai",
"max_input_tokens": 262144,

View file

@ -7,7 +7,7 @@ layer; ``litellm.types.utils`` re-exports them for backwards compatibility.
from collections.abc import Mapping
from pydantic import BaseModel, model_validator
from pydantic import BaseModel, Field, model_validator
class CredentialBase(BaseModel):
@ -17,6 +17,10 @@ class CredentialBase(BaseModel):
class CredentialItem(CredentialBase):
credential_values: dict
# PATCH-only instruction naming keys to drop from the stored credential_values. It describes an
# edit rather than the credential, so it stays out of dumps: those feed config loading, the DB
# write, and the in-memory list, none of which have a place for it.
credential_values_to_delete: tuple[str, ...] | None = Field(default=None, exclude=True)
class CreateCredentialItem(CredentialBase):
@ -36,3 +40,4 @@ class UpdateCredentialItem(BaseModel):
credential_info: Mapping[str, object]
credential_values: Mapping[str, object] | None = None
model_id: str | None = None
credential_values_to_delete: tuple[str, ...] | None = None

View file

@ -171,6 +171,42 @@ async def _session_key_is_live(session_key: str | None) -> bool:
return True
async def get_authenticated_browser_user_id(request: Request) -> str | None:
from datetime import datetime, timezone
from pydantic import TypeAdapter, ValidationError
from litellm.proxy._types import hash_token
from litellm.proxy.auth.auth_checks import ExperimentalUIJWTToken, get_key_object
from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache
user_id, session_key = _session_identity_from_cookie(request)
if not user_id or not session_key or prisma_client is None:
return None
try:
auth: Final = (
await get_key_object(
hash_token(session_key),
prisma_client,
user_api_key_cache,
proxy_logging_obj=proxy_logging_obj,
check_db_only=True,
)
if session_key.startswith("sk-")
else ExperimentalUIJWTToken.get_key_object_from_ui_hash_key(session_key)
)
except Exception:
return None
if auth is None or auth.user_id != user_id or auth.blocked or auth.expires is None:
return None
try:
expiration: Final = TypeAdapter(datetime).validate_python(auth.expires)
except ValidationError:
return None
expires: Final = expiration.replace(tzinfo=timezone.utc) if expiration.tzinfo is None else expiration
return user_id if expires > datetime.now(timezone.utc) else None
async def _byok_session_auth(request: Request) -> UserAPIKeyAuth:
"""Require the UI session cookie, with the embedded session key
re-resolved against the DB so a revoked (logged-out) session cannot

View file

@ -34885,6 +34885,10 @@
"WorkerCreated": {
"additionalProperties": false,
"properties": {
"image": {
"title": "Image",
"type": "string"
},
"token": {
"title": "Token",
"type": "string"
@ -34894,6 +34898,7 @@
}
},
"required": [
"image",
"worker",
"token"
],

View file

@ -484,7 +484,7 @@ def _is_model_cost_zero(model: str | list[str] | None, llm_router: Router | None
continue
try:
# Use router's get_model_group_info method directly for better reliability
model_group_info = llm_router.get_model_group_info(model_group=model_name)
model_group_info = llm_router.get_model_group_info(model_group=model_name, include_hidden=True)
if model_group_info is None:
# Model not found or no pricing info available

View file

@ -20,6 +20,7 @@ from litellm.constants import (
MINIMUM_CUSTOM_KEY_LENGTH,
STANDARD_CUSTOMER_ID_HEADERS,
)
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.safe_json_loads import safe_json_loads
from litellm.litellm_core_utils.url_utils import (
SSRFError,
@ -34,7 +35,8 @@ from litellm.proxy.common_utils.http_parsing_utils import extract_nested_form_me
from litellm.types.passthrough_endpoints.pass_through_endpoints import (
LITELLM_PASS_THROUGH_ENDPOINT_MARKER,
)
from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, Deployment
from litellm.types.router import CONFIGURABLE_CLIENTSIDE_AUTH_PARAMS, Deployment, server_owned_wif_fields_present
from litellm.types.router import reject_server_owned_wif_params as _reject_server_owned_wif_params
from litellm.types.utils import CustomPricingLiteLLMParams
@ -228,6 +230,63 @@ def _allow_model_level_clientside_configurable_parameters(
# ``extra_body.aws_web_identity_token``) without re-validating, so the
# banned-key check has to descend into it the same way it descends into
# ``litellm_embedding_config``.
# Re-exported from litellm.types.router, where it lives so the router can call it on a
# post-authentication merge without core importing from the proxy package.
reject_server_owned_wif_params = _reject_server_owned_wif_params
_CREDENTIAL_VALUES: Final = TypeAdapter(dict[str, object])
def reject_federated_credential_reference(body: Mapping[str, object]) -> None:
"""Raise ``ValueError`` if a request body picks the federated identity by credential name.
``load_credentials_from_list`` merges a named credential's values into the call, so a body
naming a federated credential moves the token exchange onto that credential's federation rule
and organization exactly as sending the federation fields inline would, which
``reject_server_owned_wif_params`` already refuses. Attaching one is a deployment decision, so
it is refused with no client-side opt-in to relax it, the same as the inline form.
Only credentials already loaded into memory can be resolved here, which is every credential the
proxy would resolve for the call itself: ``load_credentials_from_list`` reads the same list.
``route_manages_deployments`` names the routes this is not applied to, where choosing the
credential a deployment federates through is the point of the call.
"""
named: Final = body.get("litellm_credential_name")
if not isinstance(named, str) or not named:
return
wif_fields: Final = server_owned_wif_fields_present(
_CREDENTIAL_VALUES.validate_python(CredentialAccessor.get_credential_values(named))
)
if wif_fields:
raise ValueError(
f"Rejected Request: litellm_credential_name={named!r} names a credential configured for "
f"workload identity federation ({wif_fields[0]}), which a request body cannot choose. "
"A proxy admin attaches it to a deployment."
)
_DEPLOYMENT_MANAGEMENT_ROUTES: Final[frozenset[str]] = frozenset(
("/model/new", "/model/update", "/model/delete", "/health/test_connection")
)
_DEPLOYMENT_ID_UPDATE_ROUTE: Final = re.compile(r"^/model/[^/]+/update$")
def route_manages_deployments(route: str | None) -> bool:
"""Whether ``route`` configures a deployment instead of calling one.
These are the routes that reach ``ModelManagementAuthChecks.can_user_make_model_call``, where
a federated write is judged by ``_reject_non_admin_wif_write`` against what the write sets and
what the deployment already stores: a proxy admin goes through, anyone else is refused with a
403 naming the field. ``reject_federated_credential_reference`` runs ahead of that gate on
every route, so without this exemption a proxy admin could not attach a federated credential
to a deployment over the API or the Admin UI at all, leaving a static ``config.yaml`` entry as
the only way to configure the feature the rejection tells the caller to go configure.
"""
return route is not None and (
route in _DEPLOYMENT_MANAGEMENT_ROUTES or _DEPLOYMENT_ID_UPDATE_ROUTE.fullmatch(route) is not None
)
_NESTED_CONFIG_KEYS: Final[tuple[str, ...]] = ("litellm_embedding_config", "extra_body")
# Metadata containers that carry per-request configuration consumed by the
@ -380,12 +439,17 @@ def _check_banned_params(
general_settings: dict,
llm_router: Router | None,
model: str,
*,
manages_deployments: bool = False,
) -> None:
"""Raise ``ValueError`` if ``body`` carries a banned param without admin opt-in.
Shared between the root-level check and the nested-config check so a
new banned param only needs to be added in one place.
"""
reject_server_owned_wif_params(body)
if not manages_deployments:
reject_federated_credential_reference(body)
for param in _BANNED_REQUEST_BODY_PARAMS:
if param not in body:
continue
@ -472,7 +536,14 @@ def _reject_url_valued_fallback_target(value: str) -> None:
)
def is_request_body_safe(request_body: dict, general_settings: dict, llm_router: Router | None, model: str) -> bool:
def is_request_body_safe(
request_body: dict,
general_settings: dict,
llm_router: Router | None,
model: str,
*,
route: str | None = None,
) -> bool:
"""
Check if the request body is safe.
@ -500,25 +571,27 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router:
"""
if "model_list" in request_body:
raise ValueError("Rejected Request: model_list is not allowed in the request body.")
_check_banned_params(request_body, general_settings, llm_router, model)
manages_deployments: Final = route_manages_deployments(route)
_check_banned_params(request_body, general_settings, llm_router, model, manages_deployments=manages_deployments)
for nested_key in _NESTED_CONFIG_KEYS:
nested = _coerce_metadata_to_dict(request_body.get(nested_key))
if nested is not None:
_check_banned_params(nested, general_settings, llm_router, model)
_check_banned_params(nested, general_settings, llm_router, model, manages_deployments=manages_deployments)
for metadata_key in _NESTED_METADATA_KEYS:
metadata = _coerce_metadata_to_dict(request_body.get(metadata_key))
if metadata is not None:
_check_banned_params(metadata, general_settings, llm_router, model)
_check_banned_params(metadata, general_settings, llm_router, model, manages_deployments=manages_deployments)
if any(isinstance(key, str) and key.startswith(f"{metadata_key}[") for key in request_body):
_check_banned_params(
extract_nested_form_metadata(form_data=request_body, prefix=f"{metadata_key}["),
general_settings,
llm_router,
model,
manages_deployments=manages_deployments,
)
for target in iter_request_fallback_targets(request_body):
if isinstance(target, dict):
_check_banned_params(target, general_settings, llm_router, model)
_check_banned_params(target, general_settings, llm_router, model, manages_deployments=manages_deployments)
target_model = target.get("model")
if isinstance(target_model, str):
_reject_url_valued_fallback_target(target_model)
@ -526,6 +599,9 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router:
_reject_url_valued_fallback_target(target)
litellm_params: Final = _coerce_metadata_to_dict(request_body.get("litellm_params"))
if litellm_params is not None:
reject_server_owned_wif_params(litellm_params)
if not manages_deployments:
reject_federated_credential_reference(litellm_params)
litellm_params_metadata: Final = _coerce_metadata_to_dict(litellm_params.get("metadata"))
if litellm_params_metadata is not None:
_check_banned_params(
@ -533,6 +609,7 @@ def is_request_body_safe(request_body: dict, general_settings: dict, llm_router:
general_settings,
llm_router,
model,
manages_deployments=manages_deployments,
)
return True
@ -585,6 +662,7 @@ async def pre_db_read_auth_checks(
general_settings=general_settings,
llm_router=llm_router,
model=request_data.get("model", ""), # [TODO] use model passed in url as well (azure openai routes)
route=route,
)
# Check 3. Check if IP address is allowed

View file

@ -0,0 +1,189 @@
"""Shared helper for resolving a named Credential's values server-side.
Memory first (``litellm.credential_list``, already decrypted -- matching
``CredentialAccessor.get_credential_values``), then a DB decrypt fallback for a pod whose
in-memory list has not yet picked up a credential another pod just wrote or updated.
"""
import asyncio
from collections.abc import Mapping
from itertools import chain
from types import MappingProxyType
from typing import Final
import litellm
from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper
from litellm.proxy.utils import PrismaClient
from litellm.repositories.credentials_repository import CredentialsRepository
from litellm.router_utils.clientside_credential_handler import clientside_credential_keys
from litellm.types.router import (
GenericLiteLLMParams,
server_owned_wif_fields_named,
server_owned_wif_fields_present,
)
from litellm.types.utils import CredentialItem, LlmProviders, server_owned_wif_litellm_params
_LITELLM_PROVIDER_IDS: Final = frozenset(provider.value for provider in LlmProviders)
_FEDERATION_SURFACE_FIELDS: Final = frozenset(
(
*clientside_credential_keys,
"configurable_clientside_auth_params",
"litellm_credential_name",
*server_owned_wif_litellm_params,
)
)
def write_touches_federation_surface(incoming: Mapping[str, object] | None) -> bool:
"""Whether this write can move or re-scope the token a federated deployment mints.
Three groups of fields can. The federation parameters choose which server-side secret is read
and what the minted token is scoped to. ``litellm_credential_name`` resolves to those same
parameters by reference. ``api_key``, ``api_base``, ``base_url``, and the
``configurable_clientside_auth_params`` that let a caller override them decide where the
resulting token is sent. A write setting none of them leaves the federation configuration
exactly as the proxy admin left it, so renaming a federated deployment or changing its rpm
stays an ordinary team-admin edit.
"""
return incoming is not None and not _FEDERATION_SURFACE_FIELDS.isdisjoint(incoming.keys())
def stored_credential_provider(credential_provider: object) -> str | None:
"""The dashboard stores its display casing (``Anthropic``) on credentials it creates, so the
provider a credential names is the lowercased value when that is a litellm provider id."""
if not isinstance(credential_provider, str):
return None
lowered: Final = credential_provider.lower()
return lowered if lowered in _LITELLM_PROVIDER_IDS else None
def decrypted_or_stored(key: str, value: str) -> str:
"""The stored value decrypted, or as stored when it was never encrypted (a config.yaml value)."""
decrypted: Final = decrypt_value_helper(value=value, key=key)
return value if decrypted is None else decrypted
def _decrypted(db_credential: CredentialItem) -> CredentialItem:
"""The stored credential with every value decrypted, leaving already-plaintext values alone."""
decrypted_values: Final = MappingProxyType(
{key: decrypted_or_stored(key, value) for key, value in db_credential.credential_values.items()}
)
return CredentialItem(
credential_name=db_credential.credential_name,
credential_values=decrypted_values, # pyright: ignore[reportArgumentType] # declared dict[str, str], and pydantic copies this mapping into one on validation; LIT002 rules out building that dict here
credential_info=db_credential.credential_info,
)
async def hydrate_named_credential_authoritative(
credential_name: str,
prisma_client: PrismaClient | None,
) -> CredentialItem | None:
"""The stored credential, preferring the row over this pod's in-memory copy.
``hydrate_named_credential`` reads memory first, which is right when serving a request. A
management operation cannot: on a pod whose in-memory copy predates another pod's update, it
would export the superseded JWKS, or discover models against superseded values. Same reason
``named_credential_wif_fields`` reads both.
"""
if prisma_client is None:
return await hydrate_named_credential(credential_name, prisma_client)
db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name)
if db_credential is None:
return await hydrate_named_credential(credential_name, prisma_client)
return _decrypted(db_credential)
async def hydrate_named_credential(
credential_name: str,
prisma_client: PrismaClient | None,
) -> CredentialItem | None:
for credential in litellm.credential_list:
if credential.credential_name == credential_name:
return credential
if prisma_client is None:
return None
db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name)
if db_credential is None:
return None
return _decrypted(db_credential)
async def named_credential_wif_fields(
credential_name: str,
prisma_client: PrismaClient | None,
) -> tuple[str, ...]:
"""Federation field names a write to ``credential_name`` would touch, from memory AND the row.
Resolution reads memory first and stops there, which is right when serving a request. An
authorization decision cannot: a pod whose in-memory copy predates an admin adding federation
fields would see none and allow the write. This reads both and returns the union, so the gate
refuses whenever either side says the credential is server-owned.
"""
matching: Final = tuple(c for c in litellm.credential_list if c.credential_name == credential_name)
in_memory: Final = tuple(chain.from_iterable(server_owned_wif_fields_named(c.credential_values) for c in matching))
if prisma_client is None:
return in_memory
db_credential: Final = await CredentialsRepository(prisma_client).find_by_name(credential_name)
stored: Final = () if db_credential is None else server_owned_wif_fields_named(db_credential.credential_values)
return tuple(dict.fromkeys(in_memory + stored))
def submitted_litellm_params(params: GenericLiteLLMParams | None) -> Mapping[str, object] | None:
"""The fields a pydantic write actually set, as the mapping the federation gate reads.
Only the set fields belong here: ``GenericLiteLLMParams`` declares every federation field, so
the whole model would report every write as touching all of them.
"""
if params is None:
return None
return MappingProxyType({name: getattr(params, name, None) for name in params.model_fields_set})
async def effective_server_owned_wif_fields(
stored: Mapping[str, object] | None,
incoming: Mapping[str, object] | None,
prisma_client: PrismaClient | None,
) -> tuple[str, ...]:
"""Federation field names the deployment would carry AFTER this write.
Authorization has to read the resulting deployment, not the submitted payload. A patch that
names no federation field still lands on a deployment that has them, and a patch that only
attaches ``litellm_credential_name`` inherits whatever that credential holds.
The two sides are matched differently on purpose. ``stored`` is matched by VALUE, since it is
a full deployment and a declared-but-unset field is not a federation field it carries.
``incoming`` holds only the keys the write actually set, so an explicit null still counts as
touching the field.
"""
from_stored: Final = () if stored is None else server_owned_wif_fields_present(stored)
from_incoming: Final = () if incoming is None else server_owned_wif_fields_named(incoming.keys())
from_credential: Final = tuple(
chain.from_iterable(
await asyncio.gather(
*(
named_credential_wif_fields(credential_name, prisma_client)
for credential_name in _effective_credential_names(stored, incoming)
)
)
)
)
return tuple(dict.fromkeys(from_stored + from_incoming + from_credential))
def _effective_credential_names(
stored: Mapping[str, object] | None,
incoming: Mapping[str, object] | None,
) -> tuple[str, ...]:
"""Both the credential the deployment already carries and the one this write names.
Taking only the incoming name would let a write clear its way out: detaching a federated
credential, by sending ``litellm_credential_name: null`` alongside an api_key or api_base of
the caller's choosing, would leave nothing federated to find and the write would be allowed.
Detaching an administrator's federated credential is itself an administrator's action, so the
stored name counts whatever the write says.
"""
from_stored: Final = None if stored is None else stored.get("litellm_credential_name")
from_incoming: Final = None if incoming is None else incoming.get("litellm_credential_name")
return tuple(dict.fromkeys(name for name in (from_stored, from_incoming) if isinstance(name, str)))

View file

@ -9,26 +9,134 @@ from typing import (
cast, # noqa: TID251 # jsonify_object in proxy/utils.py is annotated with a bare dict
)
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response
from fastapi import APIRouter, Depends, HTTPException, Path, Request, Response, status
from pydantic import TypeAdapter
import litellm
from litellm._logging import verbose_proxy_logger
from litellm.litellm_core_utils.credential_accessor import CredentialAccessor
from litellm.litellm_core_utils.litellm_logging import _get_masked_values
from litellm.llms.anthropic.wif import (
ExportedJwks,
NotAnInternalIssuerCredential,
UnbuildableIdentitySource,
anthropic_internal_issuer_jwks,
)
from litellm.models.credentials import UpdateCredentialItem
from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth
from litellm.proxy._types import (
CommonProxyErrors,
LitellmUserRoles,
ProxyErrorTypes,
ProxyException,
UserAPIKeyAuth,
)
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
from litellm.proxy.common_utils.credential_hydration import (
hydrate_named_credential,
hydrate_named_credential_authoritative,
named_credential_wif_fields,
stored_credential_provider,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import encrypt_value_helper
from litellm.proxy.utils import handle_exception_on_proxy, jsonify_object
from litellm.repositories.base_repository import is_unique_violation
from litellm.repositories.credentials_repository import CredentialsRepository
from litellm.types.router import server_owned_wif_fields_named
from litellm.types.utils import CreateCredentialItem, CredentialItem
router: Final = APIRouter()
_CREDENTIAL_DICT_ADAPTER: Final = TypeAdapter(dict[str, object])
def _reject_non_admin_wif_fields(
wif_fields: tuple[str, ...],
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""A credential referenced by ``litellm_credential_name`` feeds its values into the same
workload identity federation resolution as a deployment's own ``litellm_params``. Only proxy
admins may touch a server-owned WIF field, whether they write it, drop it, or edit a stored
credential that already carries one.
"""
if not wif_fields or user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
return
raise ProxyException(
message=(
f"Only proxy admins can change {wif_fields[0]!r}, a server-owned workload identity federation parameter."
),
type=ProxyErrorTypes.auth_error.value,
code=status.HTTP_403_FORBIDDEN,
param=wif_fields[0],
)
def _incoming_wif_fields(incoming_values: Mapping[str, object], credential: UpdateCredentialItem) -> tuple[str, ...]:
"""WIF fields the request touches: the ones its values set (to any value, ``None`` included,
since the key alone is what the federation resolver reacts to), whether the caller sent them
or named a deployment through ``model_id`` for the proxy to copy them from, plus the ones it
names in ``credential_values_to_delete``, since dropping a federation field off the stored
credential breaks every deployment referencing it just as installing one would redirect them.
"""
return server_owned_wif_fields_named(incoming_values) + server_owned_wif_fields_named(
credential.credential_values_to_delete or ()
)
def _stored_wif_fields(stored_credential: CredentialItem) -> tuple[str, ...]:
return server_owned_wif_fields_named(stored_credential.credential_values)
def _reject_overlapping_credential_values(credential: UpdateCredentialItem) -> None:
overlap: Final = frozenset(credential.credential_values or ()) & frozenset(
credential.credential_values_to_delete or ()
)
if overlap:
raise HTTPException(
status_code=400,
detail=f"credential_values_to_delete overlaps credential_values for key(s): {sorted(overlap)}",
)
def _without_null_values(credential_values: Mapping[str, object]) -> dict[str, object]:
"""A null carries no credential, and the federation resolver refuses a foreign variant's field by
KEY, so a stored ``{"anthropic_issuer_url": null}`` wedges every deployment that names this
credential. ``model_dump(exclude_none=True)`` cannot do this: it drops the model's own null
fields, and ``credential_values`` is a mapping inside one of them.
"""
return {key: value for key, value in credential_values.items() if value is not None}
def _sync_in_memory_credential(credential: CredentialItem, credential_name: str, new_name: str) -> None:
"""Mirror a DB credential update into the in-memory ``credential_list`` used by request-time
resolution; a no-op if the credential isn't loaded in memory (e.g. proxy restarted since boot).
"""
existing_in_memory: CredentialItem | None = None
for cred in litellm.credential_list:
if cred.credential_name == credential_name:
existing_in_memory = cred
break
if existing_in_memory is None:
return
in_memory_values: Final = dict(existing_in_memory.credential_values or {})
if credential.credential_values:
in_memory_values.update(_without_null_values(credential.credential_values))
for key in credential.credential_values_to_delete or ():
in_memory_values.pop(key, None)
in_memory_info: Final = dict(existing_in_memory.credential_info or {})
if credential.credential_info:
in_memory_info.update(credential.credential_info)
updated_in_memory: Final = CredentialItem(
credential_name=new_name,
credential_values=in_memory_values,
credential_info=in_memory_info,
)
# Remove old entry if renamed, then use upsert_credentials to handle duplicates
if new_name != credential_name:
litellm.credential_list = [c for c in litellm.credential_list if c.credential_name != credential_name]
CredentialAccessor.upsert_credentials([updated_in_memory])
class CredentialHelperUtils:
@staticmethod
def encrypt_credential_values(credential: CredentialItem, new_encryption_key: str | None = None) -> CredentialItem:
@ -108,13 +216,17 @@ async def create_credential(
status_code=400,
detail="Credential values are required. Unable to infer credential values from model ID.",
)
_reject_non_admin_wif_fields(server_owned_wif_fields_named(credential_values), user_api_key_dict)
_reject_non_admin_wif_fields(
await named_credential_wif_fields(credential.credential_name, prisma_client), user_api_key_dict
)
processed_credential: Final = CredentialItem(
credential_name=credential.credential_name,
credential_values=_CREDENTIAL_DICT_ADAPTER.validate_python(credential_values),
credential_values=_without_null_values(_CREDENTIAL_DICT_ADAPTER.validate_python(credential_values)),
credential_info=credential.credential_info,
)
encrypted_credential: Final = CredentialHelperUtils.encrypt_credential_values(processed_credential)
credentials_dict: Final = encrypted_credential.model_dump()
credentials_dict: Final = encrypted_credential.model_dump(exclude_none=True)
credentials_dict_jsonified: Final = cast( # cast-ok: deep-copies a model_dump, so keys are str
"dict[str, object]", jsonify_object(credentials_dict)
)
@ -204,6 +316,65 @@ async def get_credential_by_name(
raise handle_exception_on_proxy(e)
@router.get(
"/credentials/{credential_name:path}/jwks",
dependencies=(Depends(user_api_key_auth),),
tags=["credential management"],
)
async def get_credential_internal_issuer_jwks(
credential_name: str = Path(..., description="The credential name, percent-decoded; may contain slashes"),
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), # noqa: B008 # FastAPI resolves the dependency from the default
):
"""
Export the public JWKS for an anthropic ``internal_issuer`` credential, so the operator can
register it on the Anthropic federation issuer from the UI. Never touches the private signing
key: only its derived public JWKS leaves this process. 404s for any other credential shape.
"""
from litellm.proxy.proxy_server import prisma_client
if user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise HTTPException(
status_code=403,
detail={"error": "Only proxy admins can export a credential's JWKS."},
)
try:
credential: Final = await hydrate_named_credential_authoritative(credential_name, prisma_client)
credential_provider: Final = (
None
if credential is None
else stored_credential_provider(credential.credential_info.get("custom_llm_provider"))
)
if credential is None or credential_provider != "anthropic":
raise HTTPException(
status_code=404,
detail={"error": f"No anthropic credential named {credential_name!r}."},
)
match anthropic_internal_issuer_jwks(credential.credential_values):
case ExportedJwks(document):
return Response(content=document, media_type="application/json")
case NotAnInternalIssuerCredential(required_param, required_value):
raise HTTPException(
status_code=404,
detail={
"error": (
f"Credential {credential_name!r} is not configured with "
f"{required_param}={required_value!r}."
)
},
)
case UnbuildableIdentitySource(message):
raise HTTPException(
status_code=400,
detail={"error": message},
)
except HTTPException:
raise
except Exception as e: # noqa: BLE001 # endpoint boundary: every failure becomes the proxy's error contract
verbose_proxy_logger.exception(e)
raise handle_exception_on_proxy(e)
@router.get(
"/credentials/by_model/{model_id}",
dependencies=[Depends(user_api_key_auth)],
@ -268,6 +439,9 @@ async def delete_credential(
status_code=500,
detail={"error": CommonProxyErrors.db_not_connected_error.value},
)
_reject_non_admin_wif_fields(
await named_credential_wif_fields(credential_name, prisma_client), user_api_key_dict
)
deleted: Final = await CredentialsRepository(prisma_client).delete_by_name(credential_name)
if deleted is None:
raise HTTPException(
@ -307,9 +481,10 @@ def update_db_credential(
# update litellm params
if encrypted_credential.credential_values:
# Encrypt any sensitive values
encrypted_params: Final = {k: v for k, v in encrypted_credential.credential_values.items()}
merged_credential.credential_values.update(_without_null_values(encrypted_credential.credential_values))
merged_credential.credential_values.update(encrypted_params)
for key in updated_patch.credential_values_to_delete or ():
merged_credential.credential_values.pop(key, None)
# update model info
if encrypted_credential.credential_info:
@ -340,6 +515,13 @@ async def update_credential(
from litellm.proxy.proxy_server import prisma_client
try:
_reject_overlapping_credential_values(credential)
incoming_values: Final = _CREDENTIAL_DICT_ADAPTER.validate_python(
_resolve_deployment_credentials(llm_router, credential.model_id)
if credential.model_id
else credential.credential_values or {}
)
_reject_non_admin_wif_fields(_incoming_wif_fields(incoming_values, credential), user_api_key_dict)
if prisma_client is None:
raise HTTPException(
status_code=500,
@ -349,18 +531,20 @@ async def update_credential(
db_credential: Final = await credentials_repository.find_by_name(credential_name)
if db_credential is None:
raise HTTPException(status_code=404, detail="Credential not found in DB.")
_reject_non_admin_wif_fields(_stored_wif_fields(db_credential), user_api_key_dict)
if credential.credential_name != credential_name:
shadowed_credential: Final = await hydrate_named_credential(credential.credential_name, prisma_client)
if shadowed_credential is not None:
_reject_non_admin_wif_fields(_stored_wif_fields(shadowed_credential), user_api_key_dict)
patch: Final = CredentialItem(
credential_name=credential.credential_name,
credential_info=_CREDENTIAL_DICT_ADAPTER.validate_python(credential.credential_info),
credential_values=_CREDENTIAL_DICT_ADAPTER.validate_python(
_resolve_deployment_credentials(llm_router, credential.model_id)
if credential.model_id
else credential.credential_values or {}
),
credential_values=incoming_values,
credential_values_to_delete=credential.credential_values_to_delete,
)
merged_credential: Final = update_db_credential(db_credential, patch)
credential_object_jsonified: Final = cast( # cast-ok: deep-copies a model_dump, so keys are str
"dict[str, object]", jsonify_object(merged_credential.model_dump())
"dict[str, object]", jsonify_object(merged_credential.model_dump(exclude_none=True))
)
await credentials_repository.update_by_name(
credential_name,
@ -371,29 +555,7 @@ async def update_credential(
)
# Sync in-memory credential_list (skip if not in memory - e.g., proxy restarted)
new_name: Final = merged_credential.credential_name
existing_in_memory: CredentialItem | None = None
for cred in litellm.credential_list:
if cred.credential_name == credential_name:
existing_in_memory = cred
break
if existing_in_memory is not None:
in_memory_values: Final = dict(existing_in_memory.credential_values or {})
if patch.credential_values:
in_memory_values.update(patch.credential_values)
in_memory_info: Final = dict(existing_in_memory.credential_info or {})
if patch.credential_info:
in_memory_info.update(patch.credential_info)
updated_in_memory: Final = CredentialItem(
credential_name=new_name,
credential_values=in_memory_values,
credential_info=in_memory_info,
)
# Remove old entry if renamed, then use upsert_credentials to handle duplicates
if new_name != credential_name:
litellm.credential_list = [c for c in litellm.credential_list if c.credential_name != credential_name]
CredentialAccessor.upsert_credentials([updated_in_memory])
_sync_in_memory_credential(patch, credential_name, merged_credential.credential_name)
return {"success": True, "message": "Credential updated successfully"}
except Exception as e:

View file

@ -26,16 +26,23 @@ from litellm.constants import (
DEFAULT_HEALTH_CHECK_PROMPT,
HEALTH_CHECK_TIMEOUT_SECONDS,
)
from litellm.litellm_core_utils.health_check_helpers import native_health_check_mode
from litellm.router_utils.auto_router_model_naming import (
StrategyRouterDependency,
classify_strategy_router_model,
strategy_router_dependencies,
)
from litellm.types.utils import secret_bearing_wif_litellm_params, server_owned_wif_litellm_params
# Provider routing fields. Allowed for proxy admins so they can see which
# region/version a deployment is checking; gated at the endpoint layer for
# non-admin callers (see _strip_admin_only_fields_from_health_result).
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS: Final = ("api_base", "api_version", "aws_bedrock_runtime_endpoint")
# Provider routing and workload identity federation fields. Allowed for proxy admins so they can
# see which region/version a deployment is checking and which identity it federates as; gated at
# the endpoint layer for non-admin callers (see _strip_admin_only_fields_from_health_result).
ADMIN_ONLY_HEALTH_DISPLAY_PARAMS: Final = (
"api_base",
"api_version",
"aws_bedrock_runtime_endpoint",
*(name for name in server_owned_wif_litellm_params if name not in secret_bearing_wif_litellm_params),
)
MINIMAL_DISPLAY_PARAMS: Final = frozenset({"model", "mode_error"})
@ -69,18 +76,29 @@ HEALTH_DISPLAY_PARAMS: Final = (
# endpoints that reject unknown fields with 400 "Unknown parameter:
# 'max_tokens'". Allow-list so new modes are safe by default.
# Per-deployment override: `model_info.health_check_supports_max_tokens`.
_MAX_TOKEN_SUPPORT_MODES: Final[frozenset[str]] = frozenset({"chat", "completion", "responses"})
_MAX_TOKEN_SUPPORT_MODES: Final[frozenset[str]] = frozenset({"chat", "completion", "responses", "anthropic_messages"})
def _resolve_health_check_mode(model_info: Mapping[str, object], litellm_params: Mapping[str, object]) -> str | None:
def _native_health_check_mode(model: str, provider_param: object) -> str | None:
try:
resolved_model, custom_llm_provider, _, _ = litellm.get_llm_provider(
model=model, custom_llm_provider=provider_param if isinstance(provider_param, str) else None
)
except Exception:
return None
return native_health_check_mode(model=resolved_model, custom_llm_provider=custom_llm_provider)
def resolve_health_check_mode(model_info: Mapping[str, object], litellm_params: Mapping[str, object]) -> str | None:
"""
Effective mode for a deployment's health-check probe.
Prefers operator-set `model_info.mode`; otherwise resolves it from the model
cost map, which understands `bedrock/` and cross-region inference-profile
prefixes (`us.`, `eu.`, `apac.`). Without this, non-chat Bedrock deployments
(e.g. embeddings) are probed as chat, so `max_tokens` is injected and the
request 400s on "extraneous key [max_tokens]".
Prefers operator-set `model_info.mode`; then the mode the provider requires for
that model family (Bedrock Mantle serves Claude ids on the Messages API only);
otherwise resolves it from the model cost map, which understands `bedrock/` and
cross-region inference-profile prefixes (`us.`, `eu.`, `apac.`). Without this,
non-chat Bedrock deployments (e.g. embeddings) are probed as chat, so
`max_tokens` is injected and the request 400s on "extraneous key [max_tokens]".
"""
explicit_mode: Final = model_info.get("mode")
if isinstance(explicit_mode, str):
@ -88,6 +106,9 @@ def _resolve_health_check_mode(model_info: Mapping[str, object], litellm_params:
model: Final = litellm_params.get("model")
if not isinstance(model, str):
return None
native_mode: Final = _native_health_check_mode(model, litellm_params.get("custom_llm_provider"))
if native_mode is not None:
return native_mode
try:
return litellm.get_model_info(model=model).get("mode")
except Exception:
@ -518,7 +539,7 @@ async def _run_model_health_check(model: dict):
if _is_strategy_router_deployment(litellm_params):
return {}
mode: Final = _resolve_health_check_mode(
mode: Final = resolve_health_check_mode(
model_info,
litellm_params, # any-ok: untyped router config dict
)
@ -768,7 +789,7 @@ def _update_litellm_params_for_health_check(model_info: dict, litellm_params: di
- updates the `voice` param with the `health_check_voice` for `audio_speech` mode if it exists Doc: https://docs.litellm.ai/docs/proxy/health#text-to-speech-models
- for Bedrock models with region routing (bedrock/region/model), strips the litellm routing prefix but preserves the model ID, and pins `custom_llm_provider` to `bedrock` (only when the deployment hasn't already set one, so an explicit `bedrock_converse` survives) so the bare model id still resolves to the provider (e.g. cross-region ids like `us.cohere.embed-v4:0`)
"""
mode: Final = _resolve_health_check_mode(
mode: Final = resolve_health_check_mode(
model_info,
litellm_params, # any-ok: untyped router config dict
)

View file

@ -12,6 +12,7 @@ from typing import Any, Final, Literal, TypedDict, cast
import fastapi
from fastapi import APIRouter, Depends, HTTPException, Request, Response, status
from pydantic import TypeAdapter
from typing_extensions import ReadOnly
import litellm
@ -41,6 +42,7 @@ from litellm.proxy.auth.auth_checks import (
)
from litellm.proxy.auth.auth_utils import (
_BANNED_REQUEST_BODY_PARAMS, # pyright: ignore[reportPrivateUsage] # one canonical list, shared with the request-body check
reject_server_owned_wif_params,
)
from litellm.proxy.auth.model_checks import get_key_models
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
@ -58,6 +60,7 @@ from litellm.proxy.health_check import (
deployments_targeted_by_name,
health_check_filter_kwargs_from_general_settings,
perform_health_check,
resolve_health_check_mode,
run_with_timeout,
)
from litellm.proxy.middleware.admission_control_middleware import (
@ -173,6 +176,24 @@ def _config_base_for_health_check(
return {key: value for key, value in config_params.items() if key not in _CONFIG_CONNECTION_FIELDS}
def _model_info_for_mode_resolution(
model_info: Mapping[str, object], stored_params: Mapping[str, object], request_params: Mapping[str, object]
) -> Mapping[str, object]:
stored_model: Final = stored_params.get("model")
if stored_model is None or request_params.get("model") in (None, stored_model):
return model_info
return {key: value for key, value in model_info.items() if key != "mode"}
def _string_mode_or_bad_request(params_mode: object) -> str | None:
if params_mode is None or isinstance(params_mode, str):
return params_mode
raise HTTPException(
status_code=status.HTTP_400_BAD_REQUEST,
detail={"error": f"litellm_params.mode must be a string, got {type(params_mode).__name__}"},
)
def get_callback_identifier(callback):
"""
Get the callback identifier string, handling both strings and objects.
@ -203,6 +224,7 @@ def get_callback_identifier(callback):
router: Final = APIRouter()
_OBJECT_MAPPING: Final = TypeAdapter(Mapping[str, object])
services = (
Literal[
"slack_budget_alerts",
@ -935,9 +957,9 @@ def _is_proxy_admin(user_api_key_dict: UserAPIKeyAuth) -> bool:
def _strip_admin_only_fields_from_health_result(result: dict) -> dict:
"""
Return a copy of the /health response with provider routing fields
(``ADMIN_ONLY_HEALTH_DISPLAY_PARAMS``) removed from each healthy/unhealthy
endpoint entry. Used to hide those fields from non-admin callers while
Return a copy of the /health response with the admin-only fields (provider routing plus the
workload identity federation params naming the identity a deployment mints as) removed from
each healthy/unhealthy endpoint entry. Used to hide those fields from non-admin callers while
still showing them which deployments they own and whether each one is
healthy. Proxy admins receive the unmodified result.
"""
@ -2033,11 +2055,16 @@ async def test_model_connection(
"rerank",
"realtime",
"responses",
"anthropic_messages",
"ocr",
]
| None = fastapi.Body(
None,
description="The mode to test the model with. If not provided, auto-detected from model capabilities.",
description=(
"The mode to test the model with. If not provided, resolved the way /health does: the deployment's "
"model_info.mode (only while the request tests the deployment's own model), then the mode the "
"provider requires for that model, then the model cost map."
),
),
litellm_params: dict = fastapi.Body(
None,
@ -2177,6 +2204,7 @@ async def test_model_connection(
"Could not find model %s in router: %s. Proceeding with request params only.", model_name, e
)
reject_server_owned_wif_params(request_litellm_params)
# Merge: config params (from proxy config) as base, request params override
litellm_params = {
**_config_base_for_health_check(
@ -2188,8 +2216,13 @@ async def test_model_connection(
}
resolved_model_info: Final = loaded_model_info if loaded_model_info is not None else model_info
probe_model_info: Final = _model_info_for_mode_resolution(
_OBJECT_MAPPING.validate_python(resolved_model_info or {}),
stored_params=_OBJECT_MAPPING.validate_python(config_litellm_params),
request_params=_OBJECT_MAPPING.validate_python(request_litellm_params),
)
litellm_params = _update_litellm_params_for_health_check(
model_info=resolved_model_info or {},
model_info=dict(probe_model_info),
litellm_params=litellm_params,
)
@ -2197,19 +2230,28 @@ async def test_model_connection(
await ModelManagementAuthChecks.can_user_make_model_call(
model_params=Deployment(
model_name="test_model",
litellm_params=LiteLLM_Params(**litellm_params),
litellm_params=LiteLLM_Params.model_validate(litellm_params),
model_info=resolved_model_info,
),
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
# The probe is a write of the caller's own params onto the stored deployment, so the
# caller's params are the incoming side: a probe that redirects a federated
# deployment's api_base is an admin's action, an unmodified probe of it is not.
incoming_params=request_litellm_params,
)
raw_params_mode: Final[object] = litellm_params.pop("mode", None)
probe_mode: Final = (
mode
or _string_mode_or_bad_request(raw_params_mode)
or resolve_health_check_mode(probe_model_info, _OBJECT_MAPPING.validate_python(litellm_params))
)
mode = mode or litellm_params.pop("mode", None)
result: Final = await run_with_timeout(
litellm.ahealth_check(
model_params=litellm_params,
mode=mode,
mode=probe_mode,
prompt="test from litellm",
input=["test from litellm"],
),
@ -2224,7 +2266,7 @@ async def test_model_connection(
"result": cleaned_result,
}
except HTTPException as e:
except (HTTPException, ProxyException) as e:
raise e
except Exception as e:
verbose_proxy_logger.debug("litellm.proxy.health_endpoints.test_model_connection(): Exception occurred - %s", e)

View file

@ -43,6 +43,7 @@ from litellm.proxy.lens.models import (
Worker,
WorkerCreated,
)
from litellm.proxy.lens.release import PROTOCOL_VERSION, release_tag, worker_image
from litellm.proxy.lens.repository import LensRepository, WriterDatabase
from litellm.proxy.lens.sources import ActivityAvailability, SourceReader, Storage, parse_execution
from litellm.proxy.lens.state import (
@ -421,9 +422,20 @@ class WorkerName(WorkerBilling):
name: str = Field(default="Lens worker", min_length=1)
def configured_worker_image() -> str:
if image := worker_image():
return image
raise HTTPException(
503,
"This LiteLLM build has no release identity. Use a published release, make lens-dev, "
"or build the gateway and worker from the same commit with the same LITELLM_RELEASE_TAG.",
)
@router.post("/workers/register", response_model=WorkerCreated)
async def register_worker(body: WorkerName, auth: Auth) -> WorkerCreated:
scope: Final = user_scope(auth, write=True)
image: Final = configured_worker_image()
await validate_key(body.analysis_key_id)
token: Final = "lens-" + secrets.token_urlsafe(40)
worker: Final = Worker(
@ -434,7 +446,7 @@ async def register_worker(body: WorkerName, auth: Auth) -> WorkerCreated:
last_seen=datetime(1970, 1, 1, tzinfo=timezone.utc),
)
await repository().save_worker(worker, hashlib.sha256(token.encode()).hexdigest())
return WorkerCreated(worker=worker, token=token)
return WorkerCreated(worker=worker, token=token, image=image)
@router.put("/workers/{worker_id}/billing-key", response_model=Worker)
@ -466,9 +478,11 @@ async def revoke_worker(worker_id: str, auth: Auth) -> bool:
@router.post("/worker/claim", response_model=Claim | None)
async def claim(worker: WorkerAuth, protocol_version: int = 1) -> Claim | None:
if protocol_version not in (2, 3):
raise HTTPException(409, "Upgrade the Lens worker using the current Connect worker command")
async def claim(worker: WorkerAuth, protocol_version: int = 1, worker_release: str = "") -> Claim | None:
image: Final = configured_worker_image()
expected: Final = release_tag()
if protocol_version != PROTOCOL_VERSION or worker_release != expected:
raise HTTPException(409, f"Upgrade the Lens worker to {image} and retry")
if worker.analysis_key_id is None:
raise HTTPException(409, "Assign an analysis key to this worker in Lens setup")
now: Final = datetime.now(timezone.utc)

View file

@ -252,6 +252,7 @@ class Worker(Record):
class WorkerCreated(Record):
image: str
worker: Worker
token: str

View file

@ -0,0 +1,35 @@
import os
from importlib.metadata import PackageNotFoundError, distribution
from pathlib import Path
from typing import Final
PROTOCOL_VERSION: Final = 4
def release_tag() -> str:
if "LITELLM_RELEASE_TAG" in os.environ:
return os.environ["LITELLM_RELEASE_TAG"]
try:
installed: Final = distribution("litellm")
except PackageNotFoundError:
return ""
if installed.read_text("direct_url.json") is not None:
return ""
if Path(str(installed.locate_file("litellm/proxy/lens/release.py"))).resolve() != Path(__file__).resolve():
return ""
from packaging.version import Version
parsed: Final = Version(installed.version)
suffix: Final = f"-dev.{parsed.dev}" if parsed.dev is not None else f"-rc.{parsed.pre[1]}" if parsed.pre else ""
return f"v{parsed.base_version}{suffix}"
def worker_image() -> str:
tag: Final = release_tag()
if not tag:
return ""
override: Final = os.environ.get("LENS_WORKER_IMAGE", "")
if override:
return override
return f"ghcr.io/berriai/litellm-lens-worker:{tag}"

View file

@ -11,6 +11,7 @@ from pydantic import BaseModel, ConfigDict, ValidationError
from .analysis import AnalysisResponseError, analyze_sample, validation_details
from .models import Claim, Coverage, ExecutionContent, ModelRequest, ModelResult, Progress, Result, Sample
from .release import PROTOCOL_VERSION, release_tag
logger: Final = logging.getLogger("litellm.lens.worker")
@ -120,8 +121,26 @@ class LensWorker:
await self.sleep(2**attempt)
return await self.model_request(path, body, attempt + 1)
async def report_unreadable_claim(self, identity: ClaimIdentity) -> None:
failure: Final = await self.client.post(
f"/lens/worker/{identity.lens_id}/{identity.job.id}/result",
json=Result(
coverage=Coverage(),
error="The worker could not read this investigation. Update the worker to match the gateway, then retry.",
).model_dump(),
)
if failure.status_code != 409:
failure.raise_for_status()
logger.warning("Worker could not read a claimed investigation; reported a version compatibility failure")
async def run_once(self) -> bool:
response: Final = await self.client.post("/lens/worker/claim", params=MappingProxyType({"protocol_version": 3}))
response: Final = await self.client.post(
"/lens/worker/claim",
params=MappingProxyType({"protocol_version": str(PROTOCOL_VERSION), "worker_release": release_tag()}),
)
if response.status_code == 409:
logger.warning("Lens worker cannot claim work: %s", response.text)
return False
response.raise_for_status()
payload: Final = response.json()
if payload is None:
@ -129,17 +148,7 @@ class LensWorker:
try:
claim: Final = Claim.model_validate(payload)
except ValidationError:
identity: Final = ClaimIdentity.model_validate(payload)
failure: Final = await self.client.post(
f"/lens/worker/{identity.lens_id}/{identity.job.id}/result",
json=Result(
coverage=Coverage(),
error="The worker could not read this investigation. Update the worker to match the gateway, then retry.",
).model_dump(),
)
if failure.status_code != 409:
failure.raise_for_status()
logger.warning("Worker could not read a claimed investigation; reported a version compatibility failure")
await self.report_unreadable_claim(ClaimIdentity.model_validate(payload))
return True
prefix: Final = f"/lens/worker/{claim.lens_id}/{claim.job.id}"

View file

@ -63,6 +63,12 @@ from litellm.proxy.common_utils.config_sync_pubsub import (
coordination_redis_cache,
publish_config_change,
)
from litellm.proxy.common_utils.credential_hydration import (
effective_server_owned_wif_fields,
hydrate_named_credential,
submitted_litellm_params,
write_touches_federation_surface,
)
from litellm.proxy.common_utils.encrypt_decrypt_utils import (
decrypt_value_helper,
encrypt_value_helper,
@ -362,6 +368,27 @@ def _raise_on_strategy_router_write_violation(
)
def _reject_non_admin_blocked_flag_on_create(
blocked: bool | None,
user_api_key_dict: UserAPIKeyAuth,
) -> None:
"""Same proxy-admin-only rule patch_model applies to the blocked flag: a team admin passed
the team-scoped auth check above, but must not be able to create a model already paused out
from under the proxy admin.
Only a blocking value is refused. A create that sends ``blocked: false`` asks for the state
every create already lands in, and dashboards and SDKs send the whole model shape on every
create, so refusing the flag's presence would turn a working non-admin create into a 403.
"""
if blocked and user_api_key_dict.user_role != LitellmUserRoles.PROXY_ADMIN:
raise ProxyException(
message="Only proxy admins can set a model's blocked flag.",
type=ProxyErrorTypes.auth_error.value,
code=status.HTTP_403_FORBIDDEN,
param="blocked",
)
def _stored_credential_name(existing_litellm_params: GenericLiteLLMParams | None) -> str | None:
if existing_litellm_params is None or existing_litellm_params.litellm_credential_name is None:
return None
@ -437,6 +464,20 @@ def _effective_complexity_router_config(
).effective
def _decrypted_litellm_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object]:
dumped: Final[Mapping[str, object]] = litellm_params.model_dump(exclude_none=True)
return MappingProxyType(
{
name: (
decrypt_value_helper(value=value, key=name, exception_type="debug", return_original_value=True)
if isinstance(value, str)
else value
)
for name, value in dumped.items()
}
)
def _effective_model(
incoming_params: GenericLiteLLMParams | None, existing_params: GenericLiteLLMParams | None
) -> str | None:
@ -1161,6 +1202,7 @@ async def patch_model(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
incoming_params=submitted_litellm_params(patch_data.litellm_params),
member_operation="update",
incoming_model_params=patch_data,
)
@ -1524,6 +1566,8 @@ async def _add_model_to_db(
}
if model_params.model_info.id is not None:
_data["model_id"] = model_params.model_info.id
if model_params.blocked is not None:
_data["blocked"] = model_params.blocked
_create_data: Final = cast("Mapping[str, object]", _data) # cast-ok: str-keyed json payload built just above
if not should_create_model_in_db:
return LiteLLM_ProxyModelTable(**_data)
@ -2105,12 +2149,51 @@ class ModelManagementAuthChecks:
)
return True
@staticmethod
async def _reject_non_admin_wif_write(
*,
model_params: Deployment,
incoming_params: Mapping[str, object] | None,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
) -> None:
if user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN:
return
if not write_touches_federation_surface(incoming_params):
return
stored: Final = _decrypted_litellm_params(model_params.litellm_params)
wif_fields: Final = await effective_server_owned_wif_fields(stored, incoming_params, prisma_client)
if wif_fields:
# ProxyException rather than HTTPException so the offending field stays a structured
# `param`, which is the contract the narrower gate this replaced already published.
raise ProxyException(
message=(
f"Only proxy admins can change the credentials of a deployment configured for "
f"workload identity federation ({wif_fields[0]!r})."
),
type=ProxyErrorTypes.auth_error.value,
code=status.HTTP_403_FORBIDDEN,
param=wif_fields[0],
)
# A name the caller expects an admin to create later would resolve to nothing today and
# start federating the moment it exists, so a non-admin may only attach one that is already there.
named: Final = None if incoming_params is None else incoming_params.get("litellm_credential_name")
if isinstance(named, str) and await hydrate_named_credential(named, prisma_client) is None:
raise ProxyException(
message=f"No credential named {named!r} exists.",
type=ProxyErrorTypes.bad_request_error.value,
code=status.HTTP_400_BAD_REQUEST,
param="litellm_credential_name",
)
@staticmethod
async def can_user_make_model_call(
model_params: Deployment,
user_api_key_dict: UserAPIKeyAuth,
prisma_client: PrismaClient,
premium_user: bool,
*,
incoming_params: Mapping[str, object] | None,
allow_missing_team: bool = False,
member_operation: Literal["create", "update"] | None = None,
incoming_model_params: updateDeployment | None = None,
@ -2120,6 +2203,19 @@ class ModelManagementAuthChecks:
LitellmUserRoles.INTERNAL_USER_VIEW_ONLY,
):
raise HTTPException(status_code=403, detail="View-only users cannot manage models.")
# Federation fields choose which server-side secret is read and where the org-scoped token
# it buys is sent, so only a proxy admin may point a federated deployment somewhere else.
# Evaluated on the RESULTING deployment: a patch attaching a credential by name inherits
# whatever that credential holds. `incoming_params` carries only the fields the write set
# and is keyword-only with no default, so a new write path cannot typecheck without
# deciding what it writes.
await ModelManagementAuthChecks._reject_non_admin_wif_write(
model_params=model_params,
incoming_params=incoming_params,
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
)
## Check team model auth
if model_params.model_info.team_id is not None:
team_obj_row: Final = await _repo_team_table(prisma_client).find_unique(
@ -2233,6 +2329,7 @@ async def delete_model(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
incoming_params=None,
allow_missing_team=True,
)
@ -2355,7 +2452,6 @@ async def delete_team_model_alias(
return removed_model_aliases
#### [BETA] - This is a beta endpoint, format might change based on user feedback. - https://github.com/BerriAI/litellm/issues/964
@router.post(
"/model/new",
description="Allows adding new models to the model list in the config.yaml",
@ -2424,10 +2520,13 @@ async def add_new_model(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
incoming_params=submitted_litellm_params(model_params.litellm_params),
member_operation="create",
)
member_write: Final = write_authorization if isinstance(write_authorization, MemberAutoRouterWrite) else None
_reject_non_admin_blocked_flag_on_create(model_params.blocked, user_api_key_dict)
ModelManagementAuthChecks.can_user_attach_credential(
litellm_params=model_params.litellm_params,
user_api_key_dict=user_api_key_dict,
@ -2625,6 +2724,7 @@ async def update_model(
user_api_key_dict=user_api_key_dict,
prisma_client=prisma_client,
premium_user=premium_user,
incoming_params=submitted_litellm_params(model_params.litellm_params),
member_operation="update",
incoming_model_params=model_params,
)

View file

@ -45,7 +45,7 @@ from litellm.constants import (
BEDROCK_AGENT_RUNTIME_PASS_THROUGH_ROUTES,
)
from litellm.litellm_core_utils.aws_partition import get_aws_dns_suffix
from litellm.llms.anthropic.common_utils import AnthropicModelInfo
from litellm.llms.anthropic.common_utils import AnthropicModelInfo, merge_anthropic_beta_headers
from litellm.llms.azure.passthrough.transformation import (
foreign_azure_deployment,
is_azure_body_model_inference_endpoint,
@ -914,7 +914,9 @@ async def anthropic_proxy_route(
is_streaming_request: Final = await is_streaming_request_fn(request)
## CREATE PASS-THROUGH
auth_header: Final = AnthropicModelInfo.get_auth_header(anthropic_api_key or None)
auth_header: Final = await AnthropicModelInfo.aget_auth_header(
anthropic_api_key or None, allow_workload_identity=True
)
endpoint_func: Final = create_pass_through_route(
endpoint=endpoint,
target=str(updated_url),
@ -2645,9 +2647,20 @@ def _upstream_headers_for_anthropic_route(
caller_headers: Final = _caller_headers_without_litellm_secrets(
request, user_api_key_dict, _HEADERS_NEVER_FORWARDED_TO_ANTHROPIC
)
if proxy_auth_header is None and _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS.isdisjoint(caller_headers):
raise HTTPException(status_code=401, detail=_CREDENTIALLESS_ANTHROPIC_MISSING_CREDENTIAL_DETAIL)
return MappingProxyType({**caller_headers, **(proxy_auth_header or {})})
if proxy_auth_header is None:
if _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS.isdisjoint(caller_headers):
raise HTTPException(status_code=401, detail=_CREDENTIALLESS_ANTHROPIC_MISSING_CREDENTIAL_DETAIL)
return caller_headers
forwarded: Final = MappingProxyType(
{name: value for name, value in caller_headers.items() if name not in _ANTHROPIC_UPSTREAM_CREDENTIAL_HEADERS}
)
caller_beta, credential_beta = caller_headers.get("anthropic-beta"), proxy_auth_header.get("anthropic-beta")
merged_beta: Final = (
{"anthropic-beta": merge_anthropic_beta_headers(caller_beta, credential_beta)}
if caller_beta and credential_beta
else {}
)
return MappingProxyType({**forwarded, **proxy_auth_header, **merged_beta})
def _upstream_headers_for_bedrock_agent_runtime_route(

View file

@ -413,6 +413,7 @@ from litellm.proxy.common_utils.callback_utils import initialize_callbacks_on_pr
from litellm.proxy.common_utils.codex_model_catalog import codex_model_list_body
from litellm.proxy.common_utils.config_includes import resolve_include_file_path, resolve_includes
from litellm.proxy.common_utils.config_sync_pubsub import ConfigSyncSubscriber
from litellm.proxy.common_utils.credential_hydration import decrypted_or_stored
from litellm.proxy.common_utils.debug_utils import init_verbose_loggers
from litellm.proxy.common_utils.debug_utils import router as debugging_endpoints_router
from litellm.proxy.common_utils.discoverable_model_filter import discoverable_rows, undiscoverable_model_names
@ -901,6 +902,7 @@ from litellm.types.router import (
RoutingGroup,
RoutingPlugin,
SearchToolTypedDict,
holds_secret_pointer,
updateDeployment,
)
from litellm.types.router import ModelInfo as RouterModelInfo
@ -5793,11 +5795,11 @@ class ProxyConfig:
return config
return {
key: self._resolved_config_value(value=value, depth=depth, max_depth=max_depth)
key: self._resolved_config_value(key=key, value=value, depth=depth, max_depth=max_depth)
for key, value in config.items()
}
def _resolved_config_value(self, value: object, depth: int, max_depth: int) -> object:
def _resolved_config_value(self, key: str, value: object, depth: int, max_depth: int) -> object:
if isinstance(value, dict):
return self._check_for_os_environ_vars(config=value, depth=depth + 1, max_depth=max_depth)
if isinstance(value, list):
@ -5807,7 +5809,7 @@ class ProxyConfig:
else item
for item in value
]
if isinstance(value, str) and value.startswith("os.environ/"):
if isinstance(value, str) and value.startswith("os.environ/") and not holds_secret_pointer(key):
resolved: Final = get_secret(value)
if resolved is None and secret_manager_would_be_consulted(value):
verbose_proxy_logger.warning("%s is absent from the configured secret manager", value)
@ -6938,7 +6940,7 @@ class ProxyConfig:
for model in model_list:
### LOAD FROM os.environ/ ###
for k, v in model["litellm_params"].items():
if isinstance(v, str) and v.startswith("os.environ/"):
if isinstance(v, str) and v.startswith("os.environ/") and not holds_secret_pointer(k):
model["litellm_params"][k] = get_secret(v)
validate_deployment_max_agentic_loops(model)
validate_deployment_complexity_router_placement(model)
@ -7349,7 +7351,7 @@ class ProxyConfig:
for model in model_list:
### LOAD FROM os.environ/ ###
for k, v in model["litellm_params"].items():
if isinstance(v, str) and v.startswith("os.environ/"):
if isinstance(v, str) and v.startswith("os.environ/") and not holds_secret_pointer(k):
model["litellm_params"][k] = get_secret(v)
## check if they have model-id's ##
@ -7394,7 +7396,11 @@ class ProxyConfig:
return value
decrypted_value: Final = decrypt_value_helper(value=value, key=key, return_original_value=True)
if isinstance(decrypted_value, str) and decrypted_value.startswith("os.environ/"):
if (
isinstance(decrypted_value, str)
and decrypted_value.startswith("os.environ/")
and not holds_secret_pointer(key)
):
return get_secret(decrypted_value)
return decrypted_value
@ -9118,7 +9124,7 @@ class ProxyConfig:
decrypted_credential_values: Final = {}
for k, v in credential_object.credential_values.items():
decrypted_credential_values[k] = decrypt_value_helper(value=v, key=k) or v
decrypted_credential_values[k] = decrypted_or_stored(k, v)
credential_object.credential_values = decrypted_credential_values
return credential_object

View file

@ -30,6 +30,8 @@ The gateway encrypts access and refresh tokens using its configured encryption k
Use **Link accounts** to associate several current or historical usernames with one internal email. Each connection has a separate username field, so a GitHub username never matches a GitLab user implicitly. Saving immediately recalculates the report without fetching repositories again. Public profile emails match automatically when they resolve unambiguously to an internal user
**Matched people only** is on by default for people, merged changes, and branch lists. Turn it off to include outside contributors and their branches. Matching depends on the linked internal account, even when no spend was recorded. This switch filters the lists; summary metrics and quality signals still cover all selected repositories
Agent-authored changes count for a person only when the supported agent metadata explicitly names a requester. Repository issue counts and revert titles are quality signals, not an individual defect score
Bug and regression counts combine repositories with issue tracking enabled. They remain unavailable when none of the selected repositories has issue tracking enabled

View file

@ -307,6 +307,7 @@ from litellm.types.router import (
RoutingStrategy,
SearchToolTypedDict,
TaggedPreRoutingStrategy,
holds_secret_pointer,
)
from litellm.types.services import ServiceTypes
from litellm.types.utils import (
@ -3987,7 +3988,7 @@ class Router:
model_info["original_model_id"] = original_model_id
deployment_pydantic_obj: Final = Deployment(
model_name=model_group,
litellm_params=LiteLLM_Params(**dynamic_litellm_params),
litellm_params=LiteLLM_Params.model_validate(dynamic_litellm_params),
model_info=model_info,
)
Router._register_deployment_pricing(deployment=deployment_pydantic_obj)
@ -9000,13 +9001,10 @@ class Router:
if access_windows_error is not None:
raise ValueError(access_windows_error)
zeroed_pricing: Final = zeroed_ptu_pricing(_model_info, _litellm_params) if config_sourced else None
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(
**( # pyright: ignore[reportArgumentType] # untyped merged dict; already true for every field here
_litellm_params
if zeroed_pricing is None
else MappingProxyType({**_litellm_params, **zeroed_pricing})
)
merged_params: Final[Mapping[str, Any]] = (
_litellm_params if zeroed_pricing is None else MappingProxyType({**_litellm_params, **zeroed_pricing})
)
litellm_params: Final[LiteLLM_Params] = LiteLLM_Params(**merged_params)
warn_on_provider_credential_mismatch(model_name=_model_name, litellm_params=_litellm_params)
deployment = Deployment(
**deployment_info,
@ -9362,7 +9360,7 @@ class Router:
continue
deployment = Deployment(
model_name=model_name,
litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params(**lp)),
litellm_params=(lp if not isinstance(lp, dict) else LiteLLM_Params.model_validate(lp)),
model_info=(entry.get("model_info") if isinstance(entry, dict) else entry.model_info),
)
if self._has_registered_strategy(self.adaptive_routers, model_name, self._deployment_tags(deployment)):
@ -9584,7 +9582,7 @@ class Router:
## check if litellm params in os.environ
if isinstance(_litellm_params, dict):
for k, v in _litellm_params.items():
if isinstance(v, str) and v.startswith("os.environ/"):
if isinstance(v, str) and v.startswith("os.environ/") and not holds_secret_pointer(k):
_litellm_params[k] = get_secret(v)
_model_info: dict = model.pop("model_info", {})
@ -10791,7 +10789,7 @@ class Router:
if isinstance(litellm_params_data, LiteLLM_Params):
litellm_params = litellm_params_data
elif isinstance(litellm_params_data, dict) and "model" in litellm_params_data:
litellm_params = LiteLLM_Params(**litellm_params_data)
litellm_params = LiteLLM_Params.model_validate(litellm_params_data)
else:
raise ValueError(
f"Deployment missing valid litellm_params. "
@ -11246,13 +11244,13 @@ class Router:
return model_group_info
def get_model_group_info(self, model_group: str) -> ModelGroupInfo | None:
def get_model_group_info(self, model_group: str, *, include_hidden: bool = False) -> ModelGroupInfo | None:
"""
For a given model group name, return the combined model info
Returns:
- ModelGroupInfo if able to construct a model group
- None if error constructing model group info or hidden model group
- None if error constructing model group info or hidden model group (unless include_hidden)
"""
## Check if model group alias
if model_group in self.model_group_alias:
@ -11260,7 +11258,7 @@ class Router:
if isinstance(item, str):
_router_model_group = item
elif isinstance(item, dict):
if item["hidden"] is True:
if item["hidden"] is True and not include_hidden:
return None
else:
_router_model_group = item["model"]
@ -12302,6 +12300,16 @@ class Router:
]
return _settings_to_return
def _switch_routing_strategy(self, routing_strategy: str | None, kwargs: Mapping[str, object]) -> None:
if routing_strategy == "lar1":
from litellm.router_strategy.lar1_routing import apply_lar1_routing_strategy
apply_lar1_routing_strategy(self, kwargs.get("routing_strategy_args"))
return
self.routing_strategy_init(
routing_strategy=routing_strategy, routing_strategy_args=kwargs.get("routing_strategy_args", {})
)
def update_settings(self, **kwargs):
"""
Update the router settings.
@ -12315,6 +12323,7 @@ class Router:
]
_existing_router_settings: Final = self.get_settings()
model_group_alias_before: Final = self.model_group_alias
rebuild_routing_groups = False
routing_args_updated = False
for var in kwargs:
@ -12338,20 +12347,7 @@ class Router:
if var == "routing_strategy":
value = self._normalize_strategy(value)
if _existing_router_settings["routing_strategy"] != value:
if value == "lar1":
from litellm.router_strategy.lar1_routing import (
apply_lar1_routing_strategy,
)
apply_lar1_routing_strategy(
self,
kwargs.get("routing_strategy_args"),
)
else:
self.routing_strategy_init(
routing_strategy=value,
routing_strategy_args=kwargs.get("routing_strategy_args", {}),
)
self._switch_routing_strategy(value, kwargs)
rebuild_routing_groups = True
elif var == "routing_strategy_args":
routing_args_updated = value != self.routing_strategy_args
@ -12362,6 +12358,9 @@ class Router:
if routing_args_updated:
self._apply_updated_routing_strategy_args()
if self.model_group_alias != model_group_alias_before:
self._invalidate_model_group_info_cache()
if rebuild_routing_groups:
routing_groups_input: Final = kwargs.get("routing_groups", self._routing_groups_input)
self._init_routing_groups(routing_groups_input)
@ -12632,7 +12631,7 @@ class Router:
if allowed_model_region is not None:
if not is_region_allowed(
litellm_params=LiteLLM_Params(**_litellm_params),
litellm_params=LiteLLM_Params.model_validate(_litellm_params),
allowed_model_region=allowed_model_region,
):
invalid_model_indices.add(idx)
@ -12650,7 +12649,7 @@ class Router:
_,
) = litellm.get_llm_provider(
model=_dep_model_for_params,
litellm_params=LiteLLM_Params(**_litellm_params),
litellm_params=LiteLLM_Params.model_validate(_litellm_params),
)
except Exception as e: # noqa: BLE001 # best-effort filter: an unresolvable provider must not fail the request
verbose_router_logger.debug(

View file

@ -13,8 +13,16 @@ Ensures cooldowns are applied correctly.
from typing import Final
from litellm.types.utils import server_owned_wif_litellm_params
clientside_credential_keys: Final = ["api_key", "api_base", "base_url"]
# Set on a deployment whose api_base was client-redirected, so the Anthropic auth path refuses to
# mint a federation token there even when WIF is configured only through ANTHROPIC_* env vars (which
# cannot be cleared from litellm_params).
DISABLE_WORKLOAD_IDENTITY_PARAM: Final = "anthropic_disable_workload_identity_federation"
_WIF_CLEAR_ON_BASE_OVERRIDE: Final = tuple(sorted(server_owned_wif_litellm_params))
def _admin_config_fields_to_clear_on_base_override() -> list[str]:
"""
@ -59,6 +67,14 @@ def _admin_config_fields_to_clear_on_base_override() -> list[str]:
# ``api_base`` for the same reason as the OCI entries above.
"nvcf_function_id",
"use_ssl",
# Workload-identity federation minting fields, restated here from
# server_owned_wif_litellm_params the same way azure_ad_token above is restated
# despite also being declared on CredentialLiteLLMParams (hence covered by
# typed_fields too): a federation token minted for a client-redirected api_base
# would send the workload's OIDC assertion, and then the minted bearer, to the
# caller-chosen host, so this list must stay correct even if a field is ever
# dropped from the typed model.
*_WIF_CLEAR_ON_BASE_OVERRIDE,
]
return typed_fields + kwargs_only_fields
@ -101,5 +117,6 @@ def get_dynamic_litellm_params(litellm_params: dict, request_kwargs: dict) -> di
litellm_params.pop(field, None)
if field in request_kwargs:
litellm_params[field] = request_kwargs[field]
litellm_params[DISABLE_WORKLOAD_IDENTITY_PARAM] = True
return litellm_params

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