mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-11 03:38:38 +00:00
ci: merge main into litellm_chore_f0413c
This commit is contained in:
commit
822a9ddbc7
246 changed files with 24814 additions and 2519 deletions
10
.github/workflows/image-scan.yml
vendored
10
.github/workflows/image-scan.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
9
.github/workflows/lens-worker.yml
vendored
9
.github/workflows/lens-worker.yml
vendored
|
|
@ -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 \
|
||||
|
|
|
|||
1
.github/workflows/test-unit.yml
vendored
1
.github/workflows/test-unit.yml
vendored
|
|
@ -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
|
||||
|
|
|
|||
13
Dockerfile
13
Dockerfile
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -22,6 +22,7 @@ BACKEND_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/customer/",
|
||||
"/end_user/",
|
||||
"/sso/",
|
||||
"/liteadmin/slack/connect/",
|
||||
"/login",
|
||||
"/v2/login",
|
||||
"/v3/login",
|
||||
|
|
|
|||
|
|
@ -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)",
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
7
deploy/lens/config.yaml
Normal 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
91
deploy/lens/stack.yaml
Normal 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:
|
||||
44
docker-compose.liteadmin.yml
Normal file
44
docker-compose.liteadmin.yml
Normal 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:
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
283
enterprise/litellm_enterprise/proxy/liteadmin.py
Normal file
283
enterprise/litellm_enterprise/proxy/liteadmin.py
Normal 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
|
||||
|
|
@ -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 }}
|
||||
|
|
|
|||
112
helm/litellm-helm/templates/liteadmin.yaml
Normal file
112
helm/litellm-helm/templates/liteadmin.yaml
Normal 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 }}
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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) -}}
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
72
helm/litellm/templates/lens/deployment.yaml
Normal file
72
helm/litellm/templates/lens/deployment.yaml
Normal 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 }}
|
||||
114
helm/litellm/tests/lens_worker_tests.yaml
Normal file
114
helm/litellm/tests/lens_worker_tests.yaml
Normal 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
|
||||
|
|
@ -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: {}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
}
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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);
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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(|| {
|
||||
|
|
|
|||
|
|
@ -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 {
|
||||
|
|
|
|||
|
|
@ -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());
|
||||
|
|
|
|||
|
|
@ -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)]
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
{
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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})
|
||||
),
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
604
litellm/llms/anthropic/wif.py
Normal file
604
litellm/llms/anthropic/wif.py
Normal 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)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
99
litellm/llms/base_llm/auth/__init__.py
Normal file
99
litellm/llms/base_llm/auth/__init__.py
Normal 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",
|
||||
)
|
||||
225
litellm/llms/base_llm/auth/client_credentials.py
Normal file
225
litellm/llms/base_llm/auth/client_credentials.py
Normal 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)
|
||||
76
litellm/llms/base_llm/auth/identity_source.py
Normal file
76
litellm/llms/base_llm/auth/identity_source.py
Normal 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>"
|
||||
86
litellm/llms/base_llm/auth/internal_issuer.py
Normal file
86
litellm/llms/base_llm/auth/internal_issuer.py
Normal 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))
|
||||
115
litellm/llms/base_llm/auth/jwt_signing.py
Normal file
115
litellm/llms/base_llm/auth/jwt_signing.py
Normal 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)
|
||||
181
litellm/llms/base_llm/auth/shared_token_store.py
Normal file
181
litellm/llms/base_llm/auth/shared_token_store.py
Normal 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()}")
|
||||
941
litellm/llms/base_llm/auth/token_exchange.py
Normal file
941
litellm/llms/base_llm/auth/token_exchange.py
Normal 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())
|
||||
100
litellm/llms/base_llm/auth/types.py
Normal file
100
litellm/llms/base_llm/auth/types.py
Normal 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: ...
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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")
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"
|
||||
],
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
189
litellm/proxy/common_utils/credential_hydration.py
Normal file
189
litellm/proxy/common_utils/credential_hydration.py
Normal 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)))
|
||||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -252,6 +252,7 @@ class Worker(Record):
|
|||
|
||||
|
||||
class WorkerCreated(Record):
|
||||
image: str
|
||||
worker: Worker
|
||||
token: str
|
||||
|
||||
|
|
|
|||
35
litellm/proxy/lens/release.py
Normal file
35
litellm/proxy/lens/release.py
Normal 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}"
|
||||
|
|
@ -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}"
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
Loading…
Add table
Reference in a new issue