diff --git a/.github/workflows/test-e2e-changed.yml b/.github/workflows/test-e2e-changed.yml index c9f08deb36e..6da16a33ea3 100644 --- a/.github/workflows/test-e2e-changed.yml +++ b/.github/workflows/test-e2e-changed.yml @@ -175,6 +175,8 @@ jobs: env: TESTS: ${{ needs.detect.outputs.tests }} E2E_FIXTURE_MODE: live + E2E_PROVIDER_EDGE_HOST_REACHABLE: '1' + COLUMNS: '400' run: | umask 077 read -r -a test_files <<< "${TESTS}" @@ -189,6 +191,7 @@ jobs: uv run --no-sync python .github/e2e-stack/assert_tests_ran.py "${report}" "${test_files[@]}" verified=$? set -e + grep -E '^(FAILED|ERROR) ' "${log}" || true grep -E '^=+ .* in [0-9.]+s( \([0-9:]+\))? =+$' "${log}" | tail -n 1 echo "::endgroup::" if [ "${status}" = "5" ]; then diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 551f783d4f9..278fa7c425f 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -120,6 +120,20 @@ jobs: - run: cargo test --workspace --doc --locked + - name: Test token counter feature combinations + run: | + for features in '' fast huggingface tiktoken fast,huggingface fast,tiktoken huggingface,tiktoken fast,huggingface,tiktoken; do + cargo test -p litellm-token-counter --locked --no-default-features --features "$features" + cargo check -p litellm-python-bridge --locked --no-default-features --features "abi3${features:+,$features}" + done + + - name: Test secret manager feature combinations + run: | + cargo test -p litellm-auth-gcp --locked --no-default-features + for features in '' aws google aws,google; do + cargo test -p litellm-secrets --locked --no-default-features --features "$features" + done + rust-wheel: runs-on: ubuntu-latest timeout-minutes: 30 diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index aa82a0bf3ee..49e6d7040d4 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -51,7 +51,7 @@ jobs: include: - shard: mcp-integration artifact-name: mcp-integration - test-path: "tests/mcp_tests" + test-path: "tests/mcp_tests tests/test_litellm/experimental_mcp_client" workers: 2 reruns: 0 timeout-minutes: 20 @@ -113,7 +113,6 @@ jobs: tests/test_litellm/compression tests/test_litellm/containers tests/test_litellm/endpoints - tests/test_litellm/experimental_mcp_client tests/test_litellm/models tests/test_litellm/repositories tests/test_litellm/images diff --git a/helm/litellm/templates/backend/deployment.yaml b/helm/litellm/templates/backend/deployment.yaml index 0db2f0b3d43..3eb64e5528c 100644 --- a/helm/litellm/templates/backend/deployment.yaml +++ b/helm/litellm/templates/backend/deployment.yaml @@ -7,6 +7,9 @@ metadata: {{- include "litellm.commonLabels" . | nindent 4 }} app.kubernetes.io/component: backend spec: + {{- if and (not .Values.backend.hpa.enabled) (not (kindIs "invalid" .Values.backend.replicaCount)) }} + replicas: {{ .Values.backend.replicaCount }} + {{- end }} {{- with .Values.backend.strategy }} strategy: {{- toYaml . | nindent 4 }} diff --git a/helm/litellm/templates/gateway/deployment.yaml b/helm/litellm/templates/gateway/deployment.yaml index c06cc9583a0..49b452b3053 100644 --- a/helm/litellm/templates/gateway/deployment.yaml +++ b/helm/litellm/templates/gateway/deployment.yaml @@ -7,6 +7,9 @@ metadata: {{- include "litellm.commonLabels" . | nindent 4 }} app.kubernetes.io/component: gateway spec: + {{- if and (not .Values.gateway.hpa.enabled) (not (kindIs "invalid" .Values.gateway.replicaCount)) }} + replicas: {{ .Values.gateway.replicaCount }} + {{- end }} {{- with .Values.gateway.strategy }} strategy: {{- toYaml . | nindent 4 }} diff --git a/helm/litellm/templates/ui/deployment.yaml b/helm/litellm/templates/ui/deployment.yaml index b992b347bad..efee2d5fc34 100644 --- a/helm/litellm/templates/ui/deployment.yaml +++ b/helm/litellm/templates/ui/deployment.yaml @@ -7,6 +7,9 @@ metadata: {{- include "litellm.commonLabels" . | nindent 4 }} app.kubernetes.io/component: ui spec: + {{- if and (not .Values.ui.hpa.enabled) (not (kindIs "invalid" .Values.ui.replicaCount)) }} + replicas: {{ .Values.ui.replicaCount }} + {{- end }} {{- with .Values.ui.strategy }} strategy: {{- toYaml . | nindent 4 }} diff --git a/helm/litellm/tests/replica_count_tests.yaml b/helm/litellm/tests/replica_count_tests.yaml new file mode 100644 index 00000000000..791e47ff798 --- /dev/null +++ b/helm/litellm/tests/replica_count_tests.yaml @@ -0,0 +1,100 @@ +suite: test fixed replica count when HPA is disabled +templates: + - gateway/deployment.yaml + - gateway/configmap.yaml + - backend/deployment.yaml + - ui/deployment.yaml +values: + - ./values/required.yaml +tests: + - it: gateway renders replicaCount into spec.replicas when its HPA is disabled + template: gateway/deployment.yaml + set: + gateway.hpa.enabled: false + gateway.replicaCount: 3 + asserts: + - isKind: + of: Deployment + - equal: + path: spec.replicas + value: 3 + + - it: backend renders replicaCount into spec.replicas when its HPA is disabled + template: backend/deployment.yaml + set: + backend.hpa.enabled: false + backend.replicaCount: 2 + asserts: + - equal: + path: spec.replicas + value: 2 + + - it: ui renders replicaCount into spec.replicas when its HPA is disabled + template: ui/deployment.yaml + set: + ui.hpa.enabled: false + ui.replicaCount: 2 + asserts: + - equal: + path: spec.replicas + value: 2 + + - it: replicaCount 0 scales the gateway to zero instead of being treated as unset + template: gateway/deployment.yaml + set: + gateway.hpa.enabled: false + gateway.replicaCount: 0 + asserts: + - equal: + path: spec.replicas + value: 0 + + - it: a component with HPA disabled but no replicaCount set keeps omitting spec.replicas, so upgrades do not reset a hand-scaled Deployment + set: + gateway.hpa.enabled: false + backend.hpa.enabled: false + ui.hpa.enabled: false + asserts: + - notExists: + path: spec.replicas + template: gateway/deployment.yaml + - notExists: + path: spec.replicas + template: backend/deployment.yaml + - notExists: + path: spec.replicas + template: ui/deployment.yaml + + - it: every component omits spec.replicas when its HPA is enabled, so the autoscaler owns the count + set: + gateway.hpa.enabled: true + gateway.replicaCount: 3 + backend.hpa.enabled: true + backend.replicaCount: 3 + ui.hpa.enabled: true + ui.replicaCount: 3 + asserts: + - notExists: + path: spec.replicas + template: gateway/deployment.yaml + - notExists: + path: spec.replicas + template: backend/deployment.yaml + - notExists: + path: spec.replicas + template: ui/deployment.yaml + + - it: a component with HPA disabled renders replicas while a sibling with HPA enabled does not + set: + gateway.hpa.enabled: false + gateway.replicaCount: 4 + backend.hpa.enabled: true + backend.replicaCount: 4 + asserts: + - equal: + path: spec.replicas + value: 4 + template: gateway/deployment.yaml + - notExists: + path: spec.replicas + template: backend/deployment.yaml diff --git a/helm/litellm/values.yaml b/helm/litellm/values.yaml index 4ca54131d6a..2c0c7151a32 100644 --- a/helm/litellm/values.yaml +++ b/helm/litellm/values.yaml @@ -397,6 +397,11 @@ gateway: # failureThreshold: 30 # periodSeconds: 10 startupProbe: {} + # Optional fixed pod count, rendered into the Deployment's spec.replicas only + # when hpa.enabled is false. Unset by default so an existing Deployment keeps + # its current count; with the HPA on, the autoscaler owns the count, e.g.: + # replicaCount: 3 + replicaCount: hpa: enabled: true minReplicas: 1 @@ -524,6 +529,8 @@ backend: strategy: {} # Optional startupProbe; same shape as gateway.startupProbe. Empty by default. startupProbe: {} + # Same semantics as gateway.replicaCount. + replicaCount: hpa: enabled: true minReplicas: 1 @@ -590,6 +597,8 @@ ui: strategy: {} # Optional startupProbe; same shape as gateway.startupProbe. Empty by default. startupProbe: {} + # Same semantics as gateway.replicaCount. + replicaCount: hpa: enabled: false minReplicas: 1 diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 2720cf01f2e..8c35a0be0b4 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -61,6 +61,12 @@ version = "1.0.14" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" +[[package]] +name = "anyhow" +version = "1.0.104" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "330a5ed07fa54e4702c9d6c4174f74427fc0ef6e214bbd677ae50a5099946470" + [[package]] name = "arc-swap" version = "1.9.2" @@ -76,6 +82,16 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "03918c3dbd7701a85c6b9887732e2921175f26c350b4563841d0958c21d57e6d" +[[package]] +name = "assert-json-diff" +version = "2.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e4f2b81832e72834d7518d8487a0396a28cc408186a2e8854c0f98011faf12" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "async-compression" version = "0.4.46" @@ -185,9 +201,9 @@ dependencies = [ [[package]] name = "aws-runtime" -version = "1.8.1" +version = "1.9.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7816e98ee912159f45d307e5ee6bfea4a335a55aee15f7f3e32f81a6f3000f1d" +checksum = "25b43ad47adc2517efe3d706559d94b97e50e80e0321b3cadbc9f77cee88adcd" dependencies = [ "aws-credential-types", "aws-sigv4", @@ -208,6 +224,58 @@ dependencies = [ "uuid", ] +[[package]] +name = "aws-sdk-kms" +version = "1.120.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6b0fe38fee2ba5b6cd24d32d08365b314ae1adea123649e061a7eb6300b6f5b" +dependencies = [ + "arc-swap", + "aws-credential-types", + "aws-runtime", + "aws-smithy-async", + "aws-smithy-http", + "aws-smithy-json", + "aws-smithy-observability", + "aws-smithy-runtime", + "aws-smithy-runtime-api", + "aws-smithy-schema", + "aws-smithy-types", + "aws-types", + "bytes", + "fastrand", + "http 0.2.12", + "http 1.4.2", + "regex-lite", + "tracing", +] + +[[package]] +name = "aws-sdk-secretsmanager" +version = "1.117.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2d32d781b34ab083e0dc54b4c68fc5e89ddf35e97d97bdcb9386d21325c14767" +dependencies = [ + "arc-swap", + "aws-credential-types", + "aws-runtime", + "aws-smithy-async", + "aws-smithy-http", + "aws-smithy-json", + "aws-smithy-observability", + "aws-smithy-runtime", + "aws-smithy-runtime-api", + "aws-smithy-schema", + "aws-smithy-types", + "aws-types", + "bytes", + "fastrand", + "http 0.2.12", + "http 1.4.2", + "regex-lite", + "tracing", +] + [[package]] name = "aws-sdk-sts" version = "1.108.0" @@ -237,9 +305,9 @@ dependencies = [ [[package]] name = "aws-sigv4" -version = "1.5.1" +version = "1.5.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "723c2234ad7511ceef63eab016b7ba6ff7c55590fefb96fa8467af014a07309f" +checksum = "31d955e76ff96acd555bf06fa0fa6d5bf9335fa84ae7c64481b20ae61d231f70" dependencies = [ "aws-credential-types", "aws-smithy-http", @@ -302,9 +370,9 @@ dependencies = [ [[package]] name = "aws-smithy-http-client" -version = "1.2.0" +version = "1.4.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "635d23afda0a6ab48d666c4d447c4873e8d1e83518a2be2093122397e50b838e" +checksum = "7bd25384a4e437aa8d8f339afad4b69e786b936a7cb10db668a7aaf66717b1a8" dependencies = [ "aws-smithy-async", "aws-smithy-runtime-api", @@ -332,9 +400,9 @@ dependencies = [ [[package]] name = "aws-smithy-json" -version = "0.63.0" +version = "0.63.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3dc65a121adb4b33729919fcfa14fa36fb33c1555a8f06bb0e2188dbfdc1d9ef" +checksum = "3385d469edbe8b60cc72002784652b5efca39178192aa9cc4b44c9875c6bdc18" dependencies = [ "aws-smithy-runtime-api", "aws-smithy-schema", @@ -362,9 +430,9 @@ dependencies = [ [[package]] name = "aws-smithy-runtime" -version = "1.12.0" +version = "1.14.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bea94a9ff8464016338c851e24b472d7131c388c88898a502e781815b2ee6045" +checksum = "3296253d3a91b3f938a3f2bcce4daebadcb4aa4228153fcec90f9d23532b4484" dependencies = [ "aws-smithy-async", "aws-smithy-http", @@ -388,9 +456,9 @@ dependencies = [ [[package]] name = "aws-smithy-runtime-api" -version = "1.13.0" +version = "1.16.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "22ed1ebe6e0a95ea84570225f5a8208dec4b8f77e61a9b0d6f51773fcb4612f0" +checksum = "6d881a7b7ad179fd6611680c9de89f716fb00ab40299a9a7b8c6913e8f7511a8" dependencies = [ "aws-smithy-async", "aws-smithy-runtime-api-macros", @@ -417,9 +485,9 @@ dependencies = [ [[package]] name = "aws-smithy-schema" -version = "0.2.0" +version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7d56e0a4e53127a632224e43633b0fe045fa9e1e3cfc68b9830f1115e103f910" +checksum = "e8f395d93304280b64b7632fea798d177e74897fe7f063416ce627cd6fa24829" dependencies = [ "aws-smithy-runtime-api", "aws-smithy-types", @@ -428,9 +496,9 @@ dependencies = [ [[package]] name = "aws-smithy-types" -version = "1.6.1" +version = "1.6.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d6dc683efb34b9e755675b37fedbe0103141e5b6df7bdc9eb6967756a8c167d8" +checksum = "0b791f3ac597193fe1d08b82366986eb1f5bc31f2ac6c194c0855276116c76cd" dependencies = [ "base64-simd", "bytes", @@ -466,9 +534,9 @@ dependencies = [ [[package]] name = "aws-types" -version = "1.4.0" +version = "1.6.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e957a6c6dbce82b7a91f44231c09273159703769f447cbe85e854dfe9cf67f86" +checksum = "209f3a6d82a6e9e5f94abbed94c7a26e1c052341002bf57a5fb5481f625896fc" dependencies = [ "aws-credential-types", "aws-smithy-async", @@ -574,6 +642,12 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" +[[package]] +name = "bitflags" +version = "1.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bef38d45163c2f1dde094a7dfd33ccf595c92905c8f8f4fdc18d06fb1037718a" + [[package]] name = "bitflags" version = "2.13.1" @@ -598,6 +672,17 @@ dependencies = [ "hybrid-array", ] +[[package]] +name = "bstr" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6bb31b46c14244e20ee9984b11bf5c992b91fb6939fea616e3512c8baecdbe5f" +dependencies = [ + "memchr", + "regex-automata", + "serde_core", +] + [[package]] name = "bumpalo" version = "3.20.3" @@ -615,6 +700,9 @@ name = "bytes" version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" +dependencies = [ + "serde", +] [[package]] name = "bytes-utils" @@ -1048,6 +1136,55 @@ version = "0.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "be1e0bca6c3637f992fc1cc7cbc52a78c1ef6db076dbf1059c4323d6a2048376" +[[package]] +name = "deadpool" +version = "0.12.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0be2b1d1d6ec8d846f05e137292d0b89133caf95ef33695424c09568bdd39b1b" +dependencies = [ + "deadpool-runtime", + "lazy_static", + "num_cpus", + "tokio", +] + +[[package]] +name = "deadpool-runtime" +version = "0.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "092966b41edc516079bdf31ec78a2e0588d1d0c08f78b91d8307215928642b2b" + +[[package]] +name = "defmt" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e2953bfe4f93bbd20cc71198842756f77d161884c99ebbabc41d80231ded88d1" +dependencies = [ + "bitflags 1.3.2", + "defmt-macros", +] + +[[package]] +name = "defmt-macros" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bad9c72e7ca2137e0dc3813245a0d282fd6daad32fd800af018306a9169b5fe8" +dependencies = [ + "defmt-parser", + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "defmt-parser" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10d60334b3b2e7c9d91ef8150abfb6fa4c1c39ebbcf4a81c2e346aad939fee3e" +dependencies = [ + "thiserror 2.0.19", +] + [[package]] name = "deranged" version = "0.5.8" @@ -1181,6 +1318,17 @@ dependencies = [ "pin-project-lite", ] +[[package]] +name = "fancy-regex" +version = "0.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72cf461f865c862bb7dc573f643dd6a2b6842f7c30b07882b56bd148cc2761b8" +dependencies = [ + "bit-set", + "regex-automata", + "regex-syntax", +] + [[package]] name = "fancy-regex" version = "0.19.2" @@ -1412,6 +1560,224 @@ version = "0.3.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e4eba85ea1d0a966a983acd07deee566e67395d2d96b6fb39e62b5a833f1eb0b" +[[package]] +name = "google-cloud-auth" +version = "1.16.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff461519b1a948200f163574be072753bcfb462a323f0eb426629d89872dd685" +dependencies = [ + "async-trait", + "base64 0.23.1", + "bytes", + "google-cloud-gax", + "hex", + "hmac", + "http 1.4.2", + "jiff", + "reqwest 0.13.5", + "rustc_version", + "rustls 0.23.42", + "rustls-pki-types", + "serde", + "serde_json", + "sha2 0.11.0", + "thiserror 2.0.19", + "time", + "tokio", + "url", +] + +[[package]] +name = "google-cloud-gax" +version = "1.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c5615cff28ee59cfe52fbb4c11b8b1e77f650296e2ea4f4c2b7757ac6b19e752" +dependencies = [ + "bytes", + "futures", + "google-cloud-rpc", + "google-cloud-wkt", + "http 1.4.2", + "pin-project", + "rand 0.10.2", + "serde", + "serde_json", + "thiserror 2.0.19", + "tokio", + "tokio-stream", +] + +[[package]] +name = "google-cloud-gax-internal" +version = "0.7.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d2766757d877a7a8ac23da9884cb0e3f10ed9b75a0ce59801ce6b19bf9d5819e" +dependencies = [ + "bytes", + "futures", + "google-cloud-auth", + "google-cloud-gax", + "google-cloud-rpc", + "google-cloud-wkt", + "h2 0.4.15", + "http 1.4.2", + "http-body 1.1.0", + "http-body-util", + "hyper 1.10.1", + "lazy_static", + "opentelemetry", + "opentelemetry-semantic-conventions", + "opentelemetry_sdk", + "percent-encoding", + "pin-project", + "prost", + "prost-types", + "reqwest 0.13.5", + "rustc_version", + "serde", + "serde_json", + "thiserror 2.0.19", + "tokio", + "tokio-stream", + "tonic", + "tonic-prost", + "tower", + "tracing", + "tracing-opentelemetry", +] + +[[package]] +name = "google-cloud-iam-v1" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5f962b40234b1531e6ef73f7558871c96e117231e962086c98804323fb8d2c82" +dependencies = [ + "async-trait", + "bytes", + "google-cloud-gax", + "google-cloud-gax-internal", + "google-cloud-type", + "google-cloud-wkt", + "serde", + "serde_json", + "serde_with", + "tracing", +] + +[[package]] +name = "google-cloud-kms-v1" +version = "1.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3d0f6eab19254d9abd98035cd54e3f2522d2c49abf9d93bf5425ee738d198c15" +dependencies = [ + "async-trait", + "bytes", + "google-cloud-gax", + "google-cloud-gax-internal", + "google-cloud-iam-v1", + "google-cloud-location", + "google-cloud-longrunning", + "google-cloud-lro", + "google-cloud-wkt", + "serde", + "serde_json", + "serde_with", + "tracing", +] + +[[package]] +name = "google-cloud-location" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "280d5acdba8fcb1232c0719ed788d85b7e362b82cbb425b7050d3ce46f075ede" +dependencies = [ + "async-trait", + "bytes", + "google-cloud-gax", + "google-cloud-gax-internal", + "google-cloud-wkt", + "serde", + "serde_json", + "serde_with", + "tracing", +] + +[[package]] +name = "google-cloud-longrunning" +version = "1.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1c0363c5389ffda2b55cd8a86eef4b19a3481a48467c91dc5d77f626a9572766" +dependencies = [ + "async-trait", + "bytes", + "google-cloud-gax", + "google-cloud-gax-internal", + "google-cloud-rpc", + "google-cloud-wkt", + "serde", + "serde_json", + "serde_with", + "tracing", +] + +[[package]] +name = "google-cloud-lro" +version = "1.10.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47af3deef75c14a2983c430898d960c765bddbcc9f9188ca0563108e9227cfe7" +dependencies = [ + "google-cloud-gax", + "google-cloud-gax-internal", + "google-cloud-longrunning", + "google-cloud-rpc", + "google-cloud-wkt", + "serde", + "tokio", + "tracing", +] + +[[package]] +name = "google-cloud-rpc" +version = "1.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e2162c08a89118130979ba261080e960e44cdcb2d6e2ab8ca9b1da245285d353" +dependencies = [ + "bytes", + "google-cloud-wkt", + "serde", + "serde_json", + "serde_with", +] + +[[package]] +name = "google-cloud-type" +version = "1.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63acc3a92a85f96bab021c3a3e29b53bbacc97651e1b524d4c2991960a63eb82" +dependencies = [ + "bytes", + "google-cloud-wkt", + "serde", + "serde_json", + "serde_with", +] + +[[package]] +name = "google-cloud-wkt" +version = "1.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7fccf98cfd5481a5f5a285181ab0c62123d7d47cd2bb7299448440649349e4e7" +dependencies = [ + "base64 0.22.1", + "bytes", + "serde", + "serde_json", + "serde_with", + "thiserror 2.0.19", + "time", + "url", +] + [[package]] name = "h2" version = "0.3.27" @@ -1479,6 +1845,12 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hermit-abi" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e17592d60ebacc7d5e169f4663c5f84f9161cc90328abcfe8456f41e4dfcb284" + [[package]] name = "hex" version = "0.4.3" @@ -1608,6 +1980,7 @@ dependencies = [ "http 1.4.2", "http-body 1.1.0", "httparse", + "httpdate", "itoa", "pin-project-lite", "smallvec", @@ -1647,6 +2020,19 @@ dependencies = [ "webpki-roots", ] +[[package]] +name = "hyper-timeout" +version = "0.5.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2b90d566bffbce6a75bd8b09a05aa8c2cb1fabb6cb348f8840c9e4c90a0d83b0" +dependencies = [ + "hyper 1.10.1", + "hyper-util", + "pin-project-lite", + "tokio", + "tower-service", +] + [[package]] name = "hyper-util" version = "0.1.20" @@ -1856,6 +2242,43 @@ version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jiff" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ab1baf72f08796de0260609515130699b890ac25f30e610ad894bc5856cafdb" +dependencies = [ + "defmt", + "jiff-core", + "jiff-static", + "log", + "portable-atomic", + "portable-atomic-util", + "serde_core", +] + +[[package]] +name = "jiff-core" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5e52fe76043ccecc9005d2305ebaadf7d7fc0cc89ca6baa10a94d6bc68c7128c" +dependencies = [ + "defmt", + "log", +] + +[[package]] +name = "jiff-static" +version = "0.2.37" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "378268a1116ad67ae6228701118ac9f491d78fda38a40a1f1a9e1348de6f7212" +dependencies = [ + "jiff-core", + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "jni" version = "0.22.4" @@ -1926,6 +2349,27 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "jsonwebtoken" +version = "11.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e75fe14a82d81e5f5af639997db37d8b96045938a7ac6ab18cdbe1c7467e05e1" +dependencies = [ + "base64 0.22.1", + "getrandom 0.2.17", + "js-sys", + "serde", + "serde_json", + "signature", + "zeroize", +] + +[[package]] +name = "lazy_static" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbd2bcb4c963f2ddae06a2efc7e9f3591312473c50c6685e1f298068316e66fe" + [[package]] name = "libc" version = "0.2.186" @@ -1942,11 +2386,10 @@ checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" name = "litellm-auth" version = "0.1.0" dependencies = [ - "serde", - "subtle", - "thiserror 2.0.19", - "tokio", - "veil", + "litellm-auth-aws", + "litellm-auth-azure", + "litellm-auth-gcp", + "litellm-auth-types", ] [[package]] @@ -1959,7 +2402,7 @@ dependencies = [ "aws-sigv4", "aws-smithy-runtime-api", "aws-types", - "litellm-auth", + "litellm-auth-types", "litellm-http", "moka", "reqwest 0.12.28", @@ -1975,7 +2418,7 @@ version = "0.1.0" dependencies = [ "azure_core", "azure_identity", - "litellm-auth", + "litellm-auth-types", "moka", "rstest", "serde_json", @@ -1990,13 +2433,26 @@ name = "litellm-auth-gcp" version = "0.1.0" dependencies = [ "gcp_auth", - "litellm-auth", + "google-cloud-auth", + "http 1.4.2", + "litellm-auth-types", "moka", "serde_json", "sha2 0.10.9", "tokio", ] +[[package]] +name = "litellm-auth-types" +version = "0.1.0" +dependencies = [ + "serde", + "subtle", + "thiserror 2.0.19", + "tokio", + "veil", +] + [[package]] name = "litellm-cache" version = "0.1.0" @@ -2083,7 +2539,7 @@ dependencies = [ name = "litellm-core-utils" version = "0.1.0" dependencies = [ - "fancy-regex", + "fancy-regex 0.19.2", "litellm-types", "rstest", "serde", @@ -2209,13 +2665,110 @@ dependencies = [ ] [[package]] -name = "litellm-token-counter" +name = "litellm-secrets" +version = "0.1.0" +dependencies = [ + "aws-sdk-kms", + "base64 0.22.1", + "google-cloud-auth", + "google-cloud-kms-v1", + "jsonwebtoken", + "litellm-core-utils", + "litellm-secrets-aws", + "litellm-secrets-google", + "litellm-secrets-types", + "moka", + "reqwest 0.12.28", + "rstest", + "serde", + "serde_json", + "strum", + "tempfile", + "thiserror 2.0.19", + "tokio", + "wiremock", +] + +[[package]] +name = "litellm-secrets-aws" +version = "0.1.0" +dependencies = [ + "aws-credential-types", + "aws-sdk-kms", + "aws-sdk-secretsmanager", + "base64 0.22.1", + "litellm-auth-aws", + "litellm-core-utils", + "litellm-secrets-types", + "rstest", + "serde_json", + "thiserror 2.0.19", + "tokio", + "tracing", + "veil", + "wiremock", +] + +[[package]] +name = "litellm-secrets-google" version = "0.1.0" dependencies = [ "base64 0.22.1", + "google-cloud-auth", + "google-cloud-gax", + "google-cloud-kms-v1", + "litellm-auth-gcp", + "litellm-auth-types", + "litellm-core-utils", + "litellm-secrets-types", + "moka", + "percent-encoding", + "reqwest 0.12.28", + "rstest", + "serde", + "serde_json", + "thiserror 2.0.19", + "tokio", + "veil", + "wiremock", +] + +[[package]] +name = "litellm-secrets-types" +version = "0.1.0" +dependencies = [ + "litellm-auth-types", + "rstest", + "serde", + "serde_json", + "thiserror 2.0.19", + "tokio", + "veil", +] + +[[package]] +name = "litellm-token-counter" +version = "0.1.0" +dependencies = [ "criterion", "indexmap 2.14.0", "itoa", + "litellm-token-counter-fast", + "litellm-token-counter-huggingface", + "litellm-token-counter-tiktoken", + "rand 0.8.7", + "rstest", + "serde", + "serde_json", + "thiserror 2.0.19", + "tokenizers", +] + +[[package]] +name = "litellm-token-counter-fast" +version = "0.1.0" +dependencies = [ + "base64 0.22.1", "rand 0.8.7", "rstest", "rustc-hash", @@ -2226,6 +2779,22 @@ dependencies = [ "unicode-normalization-alignments", ] +[[package]] +name = "litellm-token-counter-huggingface" +version = "0.1.0" +dependencies = [ + "thiserror 2.0.19", + "tokenizers", +] + +[[package]] +name = "litellm-token-counter-tiktoken" +version = "0.1.0" +dependencies = [ + "thiserror 2.0.19", + "tiktoken-rs", +] + [[package]] name = "litellm-types" version = "0.1.0" @@ -2412,6 +2981,16 @@ dependencies = [ "autocfg", ] +[[package]] +name = "num_cpus" +version = "1.17.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91df4bbde75afed763b708b7eee1e8e7651e02d97f6d5dd763e89367e957b23b" +dependencies = [ + "hermit-abi", + "libc", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -2424,7 +3003,7 @@ version = "6.5.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0cc3cbf698f9438986c11a880c90a6d04b9de27575afd28bbf45b154b6c709e2" dependencies = [ - "bitflags", + "bitflags 2.13.1", "libc", "once_cell", "onig_sys", @@ -2452,6 +3031,42 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7c87def4c32ab89d880effc9e097653c8da5d6ef28e6b539d313baaacfbafcbe" +[[package]] +name = "opentelemetry" +version = "0.32.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b0142c63252a9e054e68a4c61a5778f7b14f576274d593f8ce883d191a099682" +dependencies = [ + "futures-core", + "futures-sink", + "js-sys", + "pin-project-lite", + "thiserror 2.0.19", + "tracing", +] + +[[package]] +name = "opentelemetry-semantic-conventions" +version = "0.32.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c913ac17a6c451661ee255f4625d143e51647ae78ebd969b75e41c4442f4fe47" + +[[package]] +name = "opentelemetry_sdk" +version = "0.32.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b59f80e1ac4d5ff7a2db8fb6c80badb7f0f3f858211fba08dd9aaec750894f9" +dependencies = [ + "futures-channel", + "futures-executor", + "futures-util", + "opentelemetry", + "percent-encoding", + "portable-atomic", + "rand 0.9.5", + "thiserror 2.0.19", +] + [[package]] name = "outref" version = "0.5.2" @@ -2587,6 +3202,15 @@ version = "1.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3d20d5497ef88037a52ff98267d066e7f11fcc5e99bbfbd58a42336193aacec3" +[[package]] +name = "portable-atomic-util" +version = "0.2.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "10ab3eb7f3becc3a1cbc4f2c6f20267996cfc1a6467a873763411b136a122715" +dependencies = [ + "portable-atomic", +] + [[package]] name = "potential_utf" version = "0.1.5" @@ -2637,7 +3261,7 @@ checksum = "4b45fcc2344c680f5025fe57779faef368840d0bd1f42f216291f0dc4ace4744" dependencies = [ "bit-set", "bit-vec", - "bitflags", + "bitflags 2.13.1", "num-traits", "rand 0.9.5", "rand_chacha 0.9.0", @@ -2648,6 +3272,38 @@ dependencies = [ "unarray", ] +[[package]] +name = "prost" +version = "0.14.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "528ac67416ff8646872a3c02cad9cc4ee5dc9f9540c9b10771855c95cb2e5ae1" +dependencies = [ + "bytes", + "prost-derive", +] + +[[package]] +name = "prost-derive" +version = "0.14.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b570b25f7617e43d59005d0990ccb79e950a423952cea19671b7a876da390adf" +dependencies = [ + "anyhow", + "itertools 0.14.0", + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "prost-types" +version = "0.14.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f94967dc7688f3054c7fac87473ffae4cc4c3904800e2d9f5b857246d8963b0a" +dependencies = [ + "prost", +] + [[package]] name = "pyo3" version = "0.29.2" @@ -2974,7 +3630,7 @@ version = "0.5.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" dependencies = [ - "bitflags", + "bitflags 2.13.1", ] [[package]] @@ -3105,6 +3761,9 @@ dependencies = [ "rustls 0.23.42", "rustls-pki-types", "rustls-platform-verifier", + "serde", + "serde_json", + "serde_urlencoded", "sync_wrapper", "tokio", "tokio-rustls 0.26.4", @@ -3194,7 +3853,7 @@ version = "1.1.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "891efababe418670775f199f0d233d84843c227a0949a883ce15b37c78d6629d" dependencies = [ - "bitflags", + "bitflags 2.13.1", "errno", "libc", "linux-raw-sys", @@ -3220,6 +3879,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "3c54fcab019b409d04215d3a17cb438fd7fbf192ee61461f20f4fe18704bc138" dependencies = [ "aws-lc-rs", + "log", "once_cell", "ring", "rustls-pki-types", @@ -3387,7 +4047,7 @@ version = "3.7.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b7f4bc775c73d9a02cde8bf7b2ec4c9d12743edf609006c7facc23998404cd1d" dependencies = [ - "bitflags", + "bitflags 2.13.1", "core-foundation", "core-foundation-sys", "libc", @@ -3547,6 +4207,15 @@ dependencies = [ "digest 0.11.3", ] +[[package]] +name = "sharded-slab" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f40ca3c46823713e0d4209592e8d6e826aa57e928f09752619fc696c499637f6" +dependencies = [ + "lazy_static", +] + [[package]] name = "shlex" version = "2.0.1" @@ -3563,6 +4232,15 @@ dependencies = [ "libc", ] +[[package]] +name = "signature" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de" +dependencies = [ + "rand_core 0.6.4", +] + [[package]] name = "simd-adler32" version = "0.3.10" @@ -3794,6 +4472,30 @@ dependencies = [ "syn 3.0.0", ] +[[package]] +name = "thread_local" +version = "1.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ad99c4c6d32803332c548b1af0540b357b3f5fc0be8f6c6bfe8b2e6ae784070" +dependencies = [ + "cfg-if", +] + +[[package]] +name = "tiktoken-rs" +version = "0.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "027853bbf8c7763b77c5c595f1c271c7d536ced7d6f83452911b944621e57fc2" +dependencies = [ + "anyhow", + "base64 0.22.1", + "bstr", + "fancy-regex 0.17.0", + "lazy_static", + "regex", + "rustc-hash", +] + [[package]] name = "time" version = "0.3.53" @@ -3939,6 +4641,18 @@ dependencies = [ "tokio", ] +[[package]] +name = "tokio-stream" +version = "0.1.19" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a3d06f0b082ba57c26b79407372e57cf2a1e28124f78e9479fe80322cf53420b" +dependencies = [ + "futures-core", + "pin-project-lite", + "tokio", + "tokio-util", +] + [[package]] name = "tokio-tungstenite" version = "0.24.0" @@ -3998,6 +4712,44 @@ dependencies = [ "winnow", ] +[[package]] +name = "tonic" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac2a5518c70fa84342385732db33fb3f44bc4cc748936eb5833d2df34d6445ef" +dependencies = [ + "base64 0.22.1", + "bytes", + "http 1.4.2", + "http-body 1.1.0", + "http-body-util", + "hyper 1.10.1", + "hyper-timeout", + "hyper-util", + "percent-encoding", + "pin-project", + "rustls-native-certs", + "sync_wrapper", + "tokio", + "tokio-rustls 0.26.4", + "tokio-stream", + "tower", + "tower-layer", + "tower-service", + "tracing", +] + +[[package]] +name = "tonic-prost" +version = "0.14.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "50849f68853be452acf590cde0b146665b8d507b3b8af17261df47e02c209ea0" +dependencies = [ + "bytes", + "prost", + "tonic", +] + [[package]] name = "tower" version = "0.5.3" @@ -4006,11 +4758,15 @@ checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" dependencies = [ "futures-core", "futures-util", + "indexmap 2.14.0", "pin-project-lite", + "slab", "sync_wrapper", "tokio", + "tokio-util", "tower-layer", "tower-service", + "tracing", ] [[package]] @@ -4020,7 +4776,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" dependencies = [ "async-compression", - "bitflags", + "bitflags 2.13.1", "bytes", "futures-core", "futures-util", @@ -4077,6 +4833,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" dependencies = [ "once_cell", + "valuable", ] [[package]] @@ -4089,6 +4846,31 @@ dependencies = [ "tracing", ] +[[package]] +name = "tracing-opentelemetry" +version = "0.33.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "adbc64cba7137545b8044cb1fe9814f7aacf3c6b5f9b45be8bb5db538befdb26" +dependencies = [ + "js-sys", + "opentelemetry", + "tracing", + "tracing-core", + "tracing-subscriber", + "web-time", +] + +[[package]] +name = "tracing-subscriber" +version = "0.3.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb7f578e5945fb242538965c2d0b04418d38ec25c79d160cd279bf0731c8d319" +dependencies = [ + "sharded-slab", + "thread_local", + "tracing-core", +] + [[package]] name = "try-lock" version = "0.2.5" @@ -4258,6 +5040,12 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "valuable" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ba73ea9cf16a25df0c8caa16c51acb937d5712a8429db78a3ee29d5dcacd3a65" + [[package]] name = "veil" version = "0.3.0" @@ -4634,6 +5422,29 @@ dependencies = [ "memchr", ] +[[package]] +name = "wiremock" +version = "0.6.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "08db1edfb05d9b3c1542e521aea074442088292f00b5f28e435c714a98f85031" +dependencies = [ + "assert-json-diff", + "base64 0.22.1", + "deadpool", + "futures", + "http 1.4.2", + "http-body-util", + "hyper 1.10.1", + "hyper-util", + "log", + "once_cell", + "regex", + "serde", + "serde_json", + "tokio", + "url", +] + [[package]] name = "wit-bindgen" version = "0.57.1" @@ -4727,6 +5538,20 @@ name = "zeroize" version = "1.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" +dependencies = [ + "zeroize_derive", +] + +[[package]] +name = "zeroize_derive" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3c50655cbb0fe3fc43170059e702f1ce5e19b84cec58dc87b037a09935c2f328" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] [[package]] name = "zerotrie" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index a6185632871..2f6f5feb4ad 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -14,9 +14,14 @@ litellm-host = { path = "crates/host" } litellm-callbacks-legacy-python = { path = "crates/callbacks-legacy-python" } litellm-framing = { path = "crates/framer" } litellm-auth = { path = "crates/auth" } +litellm-auth-types = { path = "crates/auth-types" } litellm-auth-aws = { path = "crates/auth-aws" } litellm-auth-azure = { path = "crates/auth-azure" } litellm-auth-gcp = { path = "crates/auth-gcp" } +litellm-secrets = { path = "crates/secrets" } +litellm-secrets-types = { path = "crates/secrets-types" } +litellm-secrets-aws = { path = "crates/secrets-aws" } +litellm-secrets-google = { path = "crates/secrets-google" } litellm-http = { path = "crates/http" } litellm-llms = { path = "crates/llms" } litellm-types = { path = "crates/types" } @@ -24,10 +29,15 @@ litellm-core-utils = { path = "crates/core-utils" } litellm-cache = { path = "crates/cache" } litellm-cache-memory = { path = "crates/cache-memory" } litellm-token-counter = { path = "crates/token-counter" } +litellm-token-counter-fast = { path = "crates/token-counter-fast" } +litellm-token-counter-huggingface = { path = "crates/token-counter-huggingface" } +litellm-token-counter-tiktoken = { path = "crates/token-counter-tiktoken" } litellm-host-python = { path = "crates/host-python" } bytes = "1" http = "1" +google-cloud-auth = { version = "1.16.0", default-features = false } +jsonwebtoken = { version = "11.1.0", default-features = false } hyper-util = { version = "0.1.20", default-features = false, features = ["client-proxy"] } proptest = "1.7.0" pyo3 = "0.29.2" @@ -45,6 +55,8 @@ serde_with = { version = "=3.16.1", default-features = false, features = ["std", sha2 = "0.10" subtle = "2" thiserror = "2.0" +tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] } +tiktoken-rs = "0.12.0" tokio = { version = "1", features = ["rt-multi-thread", "macros", "time", "net"] } tokio-tungstenite = { version = "0.24", default-features = false, features = ["connect", "rustls-tls-native-roots"] } futures-util = { version = "0.3", default-features = false, features = ["sink", "std"] } diff --git a/litellm-rust/crates/auth-aws/Cargo.toml b/litellm-rust/crates/auth-aws/Cargo.toml index 1f27c7bc990..1a35af48574 100644 --- a/litellm-rust/crates/auth-aws/Cargo.toml +++ b/litellm-rust/crates/auth-aws/Cargo.toml @@ -6,7 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] -litellm-auth.workspace = true +litellm-auth-types.workspace = true litellm-http.workspace = true moka = { workspace = true, features = ["sync"] } diff --git a/litellm-rust/crates/auth-aws/src/constants.rs b/litellm-rust/crates/auth-aws/src/constants.rs index be215cc9016..9e7c6bfab43 100644 --- a/litellm-rust/crates/auth-aws/src/constants.rs +++ b/litellm-rust/crates/auth-aws/src/constants.rs @@ -3,6 +3,8 @@ pub const AWS_SECRET_ACCESS_KEY: &str = "AWS_SECRET_ACCESS_KEY"; pub const AWS_SESSION_TOKEN: &str = "AWS_SESSION_TOKEN"; pub const AWS_REGION_NAME: &str = "AWS_REGION_NAME"; pub const AWS_REGION: &str = "AWS_REGION"; +pub const AWS_DEFAULT_REGION: &str = "AWS_DEFAULT_REGION"; +pub const AWS_BEDROCK_RUNTIME_ENDPOINT: &str = "AWS_BEDROCK_RUNTIME_ENDPOINT"; pub const AWS_SESSION_NAME: &str = "AWS_SESSION_NAME"; pub const AWS_PROFILE_NAME: &str = "AWS_PROFILE_NAME"; pub const AWS_ROLE_NAME: &str = "AWS_ROLE_NAME"; diff --git a/litellm-rust/crates/auth-aws/src/error.rs b/litellm-rust/crates/auth-aws/src/error.rs index f80fbce456e..d4c6ae8cf6c 100644 --- a/litellm-rust/crates/auth-aws/src/error.rs +++ b/litellm-rust/crates/auth-aws/src/error.rs @@ -22,7 +22,7 @@ pub enum Error { AwsMissingWebIdentityCredentials, } -impl From for litellm_auth::Error { +impl From for litellm_auth_types::Error { fn from(error: Error) -> Self { Self::ProviderAuthentication(error.to_string()) } @@ -34,11 +34,11 @@ mod tests { #[test] fn converts_to_shared_auth_error_without_losing_context() { - let error = litellm_auth::Error::from(Error::AwsProfile("profile not found".into())); + let error = litellm_auth_types::Error::from(Error::AwsProfile("profile not found".into())); assert_eq!( error, - litellm_auth::Error::ProviderAuthentication( + litellm_auth_types::Error::ProviderAuthentication( "AWS profile credentials failed: profile not found".into() ) ); diff --git a/litellm-rust/crates/auth-azure/Cargo.toml b/litellm-rust/crates/auth-azure/Cargo.toml index 8099506d2e5..1fd9d39d113 100644 --- a/litellm-rust/crates/auth-azure/Cargo.toml +++ b/litellm-rust/crates/auth-azure/Cargo.toml @@ -6,7 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] -litellm-auth.workspace = true +litellm-auth-types.workspace = true moka.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/auth-azure/src/credential_provider_cache.rs b/litellm-rust/crates/auth-azure/src/credential_provider_cache.rs index ab9ffc719df..cd16b27f66d 100644 --- a/litellm-rust/crates/auth-azure/src/credential_provider_cache.rs +++ b/litellm-rust/crates/auth-azure/src/credential_provider_cache.rs @@ -4,7 +4,7 @@ use std::sync::Arc; use azure_core::credentials::TokenCredential; use moka::future::Cache; -use litellm_auth::Error; +use litellm_auth_types::Error; #[derive(Clone, Debug, PartialEq, Eq, Hash)] pub(crate) struct AzureCredentialProviderCacheKey { diff --git a/litellm-rust/crates/auth-azure/src/native.rs b/litellm-rust/crates/auth-azure/src/native.rs index 5f913a8ad01..d635e559641 100644 --- a/litellm-rust/crates/auth-azure/src/native.rs +++ b/litellm-rust/crates/auth-azure/src/native.rs @@ -12,8 +12,8 @@ use azure_identity::{ }; use sha2::{Digest, Sha256}; -use litellm_auth::Error; -use litellm_auth::{InputSource, ResolvedCredential, SecretValue, Sourced}; +use litellm_auth_types::Error; +use litellm_auth_types::{InputSource, ResolvedCredential, SecretValue, Sourced}; use super::credential_provider_cache::{ AzureCredentialProviderCache, AzureCredentialProviderCacheKey, @@ -484,7 +484,7 @@ mod tests { use azure_core::{Bytes, Result}; use super::{NativeAzureRequest, NativeAzureTokenAcquirer, ValidatedAzureRequest}; - use litellm_auth::{InputSource, SecretValue, Sourced}; + use litellm_auth_types::{InputSource, SecretValue, Sourced}; fn deployment(value: T) -> Sourced { Sourced::new(value, InputSource::Deployment) @@ -649,7 +649,7 @@ mod tests { assert!(matches!( error, - litellm_auth::Error::MixedAzureCredentialSources + litellm_auth_types::Error::MixedAzureCredentialSources )); } @@ -679,7 +679,10 @@ mod tests { authority, )) .unwrap_err(); - assert!(matches!(error, litellm_auth::Error::InvalidAzureAuthority)); + assert!(matches!( + error, + litellm_auth_types::Error::InvalidAzureAuthority + )); } } } diff --git a/litellm-rust/crates/auth-azure/src/resolve.rs b/litellm-rust/crates/auth-azure/src/resolve.rs index 4e18cbb89aa..4d564b6e68a 100644 --- a/litellm-rust/crates/auth-azure/src/resolve.rs +++ b/litellm-rust/crates/auth-azure/src/resolve.rs @@ -1,5 +1,5 @@ -use litellm_auth::Error; -use litellm_auth::{ +use litellm_auth_types::Error; +use litellm_auth_types::{ CredentialFileRef, CredentialLookup, CredentialRef, InputSource, ResolvedCredential, SecretValue, Sourced, TokenProviderHandle, }; @@ -451,9 +451,9 @@ mod tests { }; use crate::native::ValidatedAzureRequest; use crate::types::AzureAuthInputs; - use litellm_auth::Error; - use litellm_auth::ResolvedCredential; - use litellm_auth::{ + use litellm_auth_types::Error; + use litellm_auth_types::ResolvedCredential; + use litellm_auth_types::{ CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialRef, CredentialResolver, CredentialResolverHandle, InputSource, SecretValue, Sourced, }; @@ -661,8 +661,8 @@ mod tests { #[derive(Debug)] struct CallerToken(&'static str); - impl litellm_auth::TokenProvider for CallerToken { - fn acquire(&self) -> litellm_auth::TokenFuture<'_> { + impl litellm_auth_types::TokenProvider for CallerToken { + fn acquire(&self) -> litellm_auth_types::TokenFuture<'_> { Box::pin(async move { Ok(ResolvedCredential::AccessToken { token: SecretValue::new(self.0), @@ -675,7 +675,7 @@ mod tests { fn caller_inputs(token: &'static str) -> AzureAuthInputs { let params = json!({"azure_ad_token": "static-token"}); AzureAuthInputs { - azure_ad_token_provider: Some(litellm_auth::TokenProviderHandle::new(Arc::new( + azure_ad_token_provider: Some(litellm_auth_types::TokenProviderHandle::new(Arc::new( CallerToken(token), ))), ..AzureAuthInputs::from_optional_params(params.as_object().unwrap()).unwrap() diff --git a/litellm-rust/crates/auth-azure/src/types.rs b/litellm-rust/crates/auth-azure/src/types.rs index 87e883a6a54..d5a00f09751 100644 --- a/litellm-rust/crates/auth-azure/src/types.rs +++ b/litellm-rust/crates/auth-azure/src/types.rs @@ -1,6 +1,6 @@ use std::collections::BTreeMap; -use litellm_auth::{ +use litellm_auth_types::{ CredentialResolverHandle, Error, InputSource, SecretValue, Sourced, TokenProviderHandle, }; use serde_json::{Map, Value}; @@ -126,7 +126,7 @@ fn source_for(sources: &BTreeMap, name: &str) -> InputSourc mod tests { use std::collections::BTreeMap; - use litellm_auth::{InputSource, Sourced}; + use litellm_auth_types::{InputSource, Sourced}; use serde_json::json; use super::{AzureAuthInputs, AzureCredentialType, ConfigValue}; diff --git a/litellm-rust/crates/auth-gcp/Cargo.toml b/litellm-rust/crates/auth-gcp/Cargo.toml index f24582db13e..0c6258a193c 100644 --- a/litellm-rust/crates/auth-gcp/Cargo.toml +++ b/litellm-rust/crates/auth-gcp/Cargo.toml @@ -5,8 +5,11 @@ edition.workspace = true license.workspace = true repository.workspace = true +[features] +google-sdk = ["dep:google-cloud-auth", "dep:http"] + [dependencies] -litellm-auth.workspace = true +litellm-auth-types.workspace = true moka.workspace = true serde_json.workspace = true @@ -14,3 +17,5 @@ sha2.workspace = true tokio.workspace = true gcp_auth = "0.12.7" +google-cloud-auth = { workspace = true, optional = true } +http = { workspace = true, optional = true } diff --git a/litellm-rust/crates/auth-gcp/src/lib.rs b/litellm-rust/crates/auth-gcp/src/lib.rs index bf619fee144..8aeddae9efc 100644 --- a/litellm-rust/crates/auth-gcp/src/lib.rs +++ b/litellm-rust/crates/auth-gcp/src/lib.rs @@ -1,13 +1,18 @@ use std::{collections::BTreeMap, future::Future, path::Path, pin::Pin, sync::Arc}; use gcp_auth::{CustomServiceAccount, TokenProvider}; -use litellm_auth::{ +use litellm_auth_types::{ CredentialPlacement, Error, InputSource, SecretValue, Sourced, http::apply_credential, }; use moka::future::Cache; use serde_json::{Map, Value}; use sha2::{Digest, Sha256}; +#[cfg(feature = "google-sdk")] +mod sdk; +#[cfg(feature = "google-sdk")] +pub use sdk::GoogleCredentials; + const CLOUD_PLATFORM_SCOPE: &str = "https://www.googleapis.com/auth/cloud-platform"; const GOOGLE_OAUTH_TOKEN_ENDPOINT: &str = "https://oauth2.googleapis.com/token"; const GOOGLE_APPLICATION_CREDENTIALS_ENV: &str = "GOOGLE_APPLICATION_CREDENTIALS"; @@ -26,19 +31,31 @@ pub struct VertexConfig { } impl VertexConfig { + pub fn new( + credentials: Option>, + project_id: Option, + location: Option, + ) -> Self { + Self { + credentials: credentials.filter(|value| !value.value().expose().trim().is_empty()), + project_id: project_id.filter(|value| !value.trim().is_empty()), + location: location.filter(|value| !value.trim().is_empty()), + } + } + pub fn from_sourced_optional_params( params: &Map, sources: &BTreeMap, ) -> Result { - Ok(Self { - credentials: optional_credentials( + Ok(Self::new( + optional_credentials( params, sources, &["vertex_credentials", "vertex_ai_credentials"], )?, - project_id: optional_string(params, &["vertex_project", "vertex_ai_project"])?, - location: optional_string(params, &["vertex_location", "vertex_ai_location"])?, - }) + optional_string(params, &["vertex_project", "vertex_ai_project"])?, + optional_string(params, &["vertex_location", "vertex_ai_location"])?, + )) } pub fn or_configured(self, project_id: Option<&str>, location: Option<&str>) -> Self { @@ -469,6 +486,39 @@ mod tests { assert_eq!(config.location(), Some("alias-location")); } + #[test] + fn typed_config_preserves_source_and_empty_value_fallback() { + let configured = VertexConfig::new( + Some(Sourced::new( + SecretValue::new("inline-json"), + InputSource::Request, + )), + Some("project".into()), + Some("location".into()), + ); + assert!(matches!( + credential_source(&configured, &|_| Some("environment-json".into())), + CredentialSource::Inline(value) if value.expose() == "inline-json" + )); + let empty = VertexConfig::new( + Some(Sourced::new(SecretValue::new(" "), InputSource::Request)), + Some(" ".into()), + Some(" ".into()), + ); + assert!(matches!( + credential_source(&empty, &|_| None), + CredentialSource::Adc + )); + assert_eq!( + get_vertex_ai_project(&empty, &|_| Some("env-project".into())).as_deref(), + Some("env-project") + ); + assert_eq!( + get_vertex_ai_location(&empty, &|_| Some("env-location".into())).as_deref(), + Some("env-location") + ); + } + #[test] fn project_and_location_prefer_input_then_environment() { let configured = diff --git a/litellm-rust/crates/auth-gcp/src/sdk.rs b/litellm-rust/crates/auth-gcp/src/sdk.rs new file mode 100644 index 00000000000..566df291b7d --- /dev/null +++ b/litellm-rust/crates/auth-gcp/src/sdk.rs @@ -0,0 +1,106 @@ +use std::sync::Arc; + +use google_cloud_auth::credentials::{CacheableResource, CredentialsProvider, EntityTag}; +use google_cloud_auth::errors::CredentialsError; +use http::{Extensions, HeaderMap, HeaderName, HeaderValue}; +use litellm_auth_types::Error; + +use crate::{VertexAuth, VertexConfig}; + +type EnvironmentLookup = dyn Fn(&str) -> Option + Send + Sync; + +pub struct GoogleCredentials { + auth: VertexAuth, + config: VertexConfig, + environment: Arc, +} + +impl GoogleCredentials { + pub fn new(config: VertexConfig, environment: Arc) -> Self { + Self { + auth: VertexAuth::default(), + config, + environment, + } + } + + pub async fn request_headers(&self) -> Result { + let response = self + .auth + .validate_environment(Vec::new(), None, &self.config, &|name| { + (self.environment)(name) + }) + .await?; + response + .headers + .into_iter() + .map(|(key, value)| { + let name = + HeaderName::from_bytes(key.as_bytes()).map_err(|_| Error::InvalidHeader)?; + let value = HeaderValue::from_str(&value).map_err(|_| Error::InvalidHeader)?; + Ok((name, value)) + }) + .collect() + } +} + +impl CredentialsProvider for GoogleCredentials { + async fn headers( + &self, + _: Extensions, + ) -> Result, CredentialsError> { + self.request_headers() + .await + .map(|data| CacheableResource::New { + entity_tag: EntityTag::new(), + data, + }) + .map_err(|_| CredentialsError::from_msg(false, "Google authentication failed")) + } + + async fn universe_domain(&self) -> Option { + None + } +} + +impl std::fmt::Debug for GoogleCredentials { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("GoogleCredentials").finish_non_exhaustive() + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn sdk_and_http_credentials_share_token_resolution_and_redaction() { + let credentials = GoogleCredentials::new( + VertexConfig::new(None, Some("project".into()), None), + Arc::new(|name| (name == "VERTEX_AI_API_KEY").then(|| "private-token".into())), + ); + let direct = credentials.request_headers().await.unwrap(); + let CacheableResource::New { data, .. } = + credentials.headers(Extensions::new()).await.unwrap() + else { + panic!("first request did not return headers"); + }; + assert_eq!(direct, data); + assert_eq!(data[http::header::AUTHORIZATION], "Bearer private-token"); + assert!(!format!("{credentials:?}").contains("private-token")); + } + + #[tokio::test] + async fn invalid_token_headers_return_a_redacted_sdk_error() { + let credentials = GoogleCredentials::new( + VertexConfig::new(None, Some("project".into()), None), + Arc::new(|name| (name == "VERTEX_AI_API_KEY").then(|| "private\nvalue".into())), + ); + assert_eq!( + credentials.request_headers().await.unwrap_err(), + Error::InvalidHeader + ); + let error = credentials.headers(Extensions::new()).await.unwrap_err(); + assert!(!format!("{error:?}").contains("private")); + } +} diff --git a/litellm-rust/crates/auth-types/Cargo.toml b/litellm-rust/crates/auth-types/Cargo.toml new file mode 100644 index 00000000000..cd65412127d --- /dev/null +++ b/litellm-rust/crates/auth-types/Cargo.toml @@ -0,0 +1,15 @@ +[package] +name = "litellm-auth-types" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +serde.workspace = true +subtle.workspace = true +thiserror.workspace = true +veil.workspace = true + +[dev-dependencies] +tokio.workspace = true diff --git a/litellm-rust/crates/auth/src/credential.rs b/litellm-rust/crates/auth-types/src/credential.rs similarity index 98% rename from litellm-rust/crates/auth/src/credential.rs rename to litellm-rust/crates/auth-types/src/credential.rs index 8ed1867622a..a5d14a43b71 100644 --- a/litellm-rust/crates/auth/src/credential.rs +++ b/litellm-rust/crates/auth-types/src/credential.rs @@ -5,9 +5,7 @@ use std::sync::Arc; use veil::Redact; -use crate::Error; - -use super::{ResolvedCredential, SecretValue, TokenProviderHandle}; +use crate::{Error, ResolvedCredential, SecretValue, TokenProviderHandle}; #[derive(Clone, Debug, PartialEq, Eq)] pub enum CredentialFileRef { diff --git a/litellm-rust/crates/auth/src/error.rs b/litellm-rust/crates/auth-types/src/error.rs similarity index 100% rename from litellm-rust/crates/auth/src/error.rs rename to litellm-rust/crates/auth-types/src/error.rs diff --git a/litellm-rust/crates/auth/src/http.rs b/litellm-rust/crates/auth-types/src/http.rs similarity index 92% rename from litellm-rust/crates/auth/src/http.rs rename to litellm-rust/crates/auth-types/src/http.rs index dd87d00e70f..0cb5839f965 100644 --- a/litellm-rust/crates/auth/src/http.rs +++ b/litellm-rust/crates/auth-types/src/http.rs @@ -40,9 +40,6 @@ pub fn apply_credential( ) } -/// How the upstream call is authenticated. API-key strategies become headers -/// in `prepare`; SigV4 covers the serialized body, so it is applied where the -/// outbound request is built. #[derive(Clone, Debug, PartialEq, Eq)] pub enum RequestAuth { Header { diff --git a/litellm-rust/crates/auth-types/src/lib.rs b/litellm-rust/crates/auth-types/src/lib.rs new file mode 100644 index 00000000000..9d399249c05 --- /dev/null +++ b/litellm-rust/crates/auth-types/src/lib.rs @@ -0,0 +1,57 @@ +#![forbid(unsafe_code)] + +mod credential; +mod error; +pub mod http; +mod policy; +mod secret; +mod token; + +use serde::{Deserialize, Serialize}; + +#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, Hash, PartialEq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum InputSource { + Request, + #[default] + Deployment, + Environment, +} + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct Sourced { + value: T, + source: InputSource, +} + +impl Sourced { + pub fn new(value: T, source: InputSource) -> Self { + Self { value, source } + } + + pub fn value(&self) -> &T { + &self.value + } + + pub fn source(&self) -> InputSource { + self.source + } + + pub fn into_value(self) -> T { + self.value + } + + pub fn map(self, map: impl FnOnce(T) -> U) -> Sourced { + Sourced::new(map(self.value), self.source) + } +} + +pub use credential::{ + CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialPlan, + CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle, +}; +pub use error::Error; +pub use http::{CredentialPlacement, RequestAuth}; +pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy}; +pub use secret::SecretValue; +pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; diff --git a/litellm-rust/crates/auth/src/policy.rs b/litellm-rust/crates/auth-types/src/policy.rs similarity index 96% rename from litellm-rust/crates/auth/src/policy.rs rename to litellm-rust/crates/auth-types/src/policy.rs index 4a1f5eeecf9..4c5c0365f0b 100644 --- a/litellm-rust/crates/auth/src/policy.rs +++ b/litellm-rust/crates/auth-types/src/policy.rs @@ -1,7 +1,5 @@ -use crate::Error; - -use super::http::apply_credential; -use super::{CredentialPlacement, ResolvedCredential}; +use crate::http::apply_credential; +use crate::{CredentialPlacement, Error, ResolvedCredential}; #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub enum CredentialPlanKind { diff --git a/litellm-rust/crates/auth/src/secret.rs b/litellm-rust/crates/auth-types/src/secret.rs similarity index 100% rename from litellm-rust/crates/auth/src/secret.rs rename to litellm-rust/crates/auth-types/src/secret.rs diff --git a/litellm-rust/crates/auth/src/token.rs b/litellm-rust/crates/auth-types/src/token.rs similarity index 95% rename from litellm-rust/crates/auth/src/token.rs rename to litellm-rust/crates/auth-types/src/token.rs index 94da5f259fb..4175641ce10 100644 --- a/litellm-rust/crates/auth/src/token.rs +++ b/litellm-rust/crates/auth-types/src/token.rs @@ -5,9 +5,7 @@ use std::time::SystemTime; use veil::Redact; -use crate::Error; - -use super::secret::SecretValue; +use crate::{Error, SecretValue}; #[derive(Clone, Debug, PartialEq, Eq)] pub enum ResolvedCredential { diff --git a/litellm-rust/crates/auth/Cargo.toml b/litellm-rust/crates/auth/Cargo.toml index 128a05c1a25..ee4900ebcc2 100644 --- a/litellm-rust/crates/auth/Cargo.toml +++ b/litellm-rust/crates/auth/Cargo.toml @@ -5,11 +5,14 @@ edition.workspace = true license.workspace = true repository.workspace = true -[dependencies] -serde.workspace = true -subtle.workspace = true -thiserror.workspace = true -veil.workspace = true +[features] +default = [] +aws = ["dep:litellm-auth-aws"] +azure = ["dep:litellm-auth-azure"] +gcp = ["dep:litellm-auth-gcp"] -[dev-dependencies] -tokio.workspace = true +[dependencies] +litellm-auth-types.workspace = true +litellm-auth-aws = { workspace = true, optional = true } +litellm-auth-azure = { workspace = true, optional = true } +litellm-auth-gcp = { workspace = true, optional = true } diff --git a/litellm-rust/crates/auth/src/lib.rs b/litellm-rust/crates/auth/src/lib.rs index c8d73c239b0..622a5b2d58b 100644 --- a/litellm-rust/crates/auth/src/lib.rs +++ b/litellm-rust/crates/auth/src/lib.rs @@ -1,55 +1,10 @@ -mod credential; -mod error; -pub mod http; -mod policy; -mod secret; -mod token; +#![forbid(unsafe_code)] -use serde::{Deserialize, Serialize}; +pub use litellm_auth_types::*; -#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, Hash, PartialEq, Serialize)] -#[serde(rename_all = "snake_case")] -pub enum InputSource { - Request, - #[default] - Deployment, - Environment, -} - -#[derive(Clone, Debug, Default, Eq, PartialEq)] -pub struct Sourced { - value: T, - source: InputSource, -} - -impl Sourced { - pub fn new(value: T, source: InputSource) -> Self { - Self { value, source } - } - - pub fn value(&self) -> &T { - &self.value - } - - pub fn source(&self) -> InputSource { - self.source - } - - pub fn into_value(self) -> T { - self.value - } - - pub fn map(self, map: impl FnOnce(T) -> U) -> Sourced { - Sourced::new(map(self.value), self.source) - } -} - -pub use credential::{ - CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialPlan, - CredentialPlanResolution, CredentialRef, CredentialResolver, CredentialResolverHandle, -}; -pub use error::Error; -pub use http::{CredentialPlacement, RequestAuth}; -pub use policy::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy}; -pub use secret::SecretValue; -pub use token::{ResolvedCredential, TokenFuture, TokenProvider, TokenProviderHandle}; +#[cfg(feature = "aws")] +pub use litellm_auth_aws as aws; +#[cfg(feature = "azure")] +pub use litellm_auth_azure as azure; +#[cfg(feature = "gcp")] +pub use litellm_auth_gcp as gcp; diff --git a/litellm-rust/crates/auth/tests/facade.rs b/litellm-rust/crates/auth/tests/facade.rs new file mode 100644 index 00000000000..f1092b15def --- /dev/null +++ b/litellm-rust/crates/auth/tests/facade.rs @@ -0,0 +1,33 @@ +use litellm_auth::{ + CredentialPlacement, CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, + ProviderAuthPolicy, ResolvedCredential, SecretValue, +}; + +const RULES: &[CredentialRule] = &[CredentialRule { + kind: CredentialPlanKind::Static, + placement: CredentialPlacement::Header("x-api-key"), +}]; + +#[test] +fn facade_applies_shared_auth_policy() { + let policy = ProviderAuthPolicy { + rules: RULES, + accepted_existing_headers: &["x-api-key"], + existing_header_behavior: ExistingHeaderBehavior::Preserve, + scope: None, + audience: None, + }; + + let headers = policy + .apply( + Vec::new(), + CredentialPlanKind::Static, + &ResolvedCredential::Static(SecretValue::new("secret")), + ) + .expect("facade policy applies"); + + assert_eq!( + headers, + vec![("x-api-key".to_string(), "secret".to_string())] + ); +} diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 3a4a579efa3..a76b069935f 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -10,10 +10,13 @@ name = "_native" crate-type = ["cdylib"] [features] -default = ["abi3"] +default = ["abi3", "fast"] abi3 = ["pyo3/abi3-py310"] extension-module = ["pyo3/extension-module"] panic-test = [] +fast = ["litellm-token-counter/fast"] +huggingface = ["litellm-token-counter/huggingface"] +tiktoken = ["litellm-token-counter/tiktoken"] [dependencies] bytes.workspace = true @@ -26,7 +29,7 @@ litellm-http.workspace = true litellm-llms.workspace = true litellm-types.workspace = true litellm-host-python.workspace = true -litellm-token-counter.workspace = true +litellm-token-counter = { path = "../token-counter", default-features = false } pyo3.workspace = true pyo3-async-runtimes.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/token_counter.rs b/litellm-rust/crates/python-bridge/src/token_counter.rs index 7dc86b78ad6..244401e6696 100644 --- a/litellm-rust/crates/python-bridge/src/token_counter.rs +++ b/litellm-rust/crates/python-bridge/src/token_counter.rs @@ -1,6 +1,11 @@ -use std::{num::NonZero, sync::Arc, thread::available_parallelism}; +use std::sync::Arc; -use litellm_host_python::{release_gil, run_async}; +#[cfg(any(feature = "fast", feature = "huggingface", feature = "tiktoken"))] +use std::{num::NonZero, thread::available_parallelism}; + +#[cfg(any(feature = "fast", feature = "huggingface", feature = "tiktoken"))] +use litellm_host_python::release_gil; +use litellm_host_python::run_async; use litellm_token_counter::{ CountableRequest, Error, InputTokenCount, TokenCounter as CoreTokenCounter, }; @@ -28,17 +33,66 @@ pub(crate) struct TokenCounter { impl TokenCounter { #[new] fn new(py: Python<'_>, tokenizer_json: &str) -> PyResult { - Self::load(py, || CoreTokenCounter::from_json(tokenizer_json)) + #[cfg(feature = "fast")] + { + Self::load(py, || CoreTokenCounter::from_json_fast(tokenizer_json)) + } + #[cfg(all(not(feature = "fast"), feature = "huggingface"))] + { + Self::load(py, || CoreTokenCounter::from_json(tokenizer_json)) + } + #[cfg(not(any(feature = "fast", feature = "huggingface")))] + { + let _ = (py, tokenizer_json); + Err(RustBridgeDeclined::new_err( + "tokenizer backend requires the fast or huggingface feature", + )) + } } #[staticmethod] fn from_cl100k_ranks(py: Python<'_>, rank_file: &str) -> PyResult { - Self::load(py, || CoreTokenCounter::from_cl100k_ranks(rank_file)) + #[cfg(feature = "fast")] + { + Self::load(py, || CoreTokenCounter::from_cl100k_ranks(rank_file)) + } + #[cfg(not(feature = "fast"))] + { + let _ = (py, rank_file); + Err(RustBridgeDeclined::new_err( + "tokenizer backend requires the fast feature", + )) + } } #[staticmethod] fn from_o200k_ranks(py: Python<'_>, rank_file: &str) -> PyResult { - Self::load(py, || CoreTokenCounter::from_o200k_ranks(rank_file)) + #[cfg(feature = "fast")] + { + Self::load(py, || CoreTokenCounter::from_o200k_ranks(rank_file)) + } + #[cfg(not(feature = "fast"))] + { + let _ = (py, rank_file); + Err(RustBridgeDeclined::new_err( + "tokenizer backend requires the fast feature", + )) + } + } + + #[staticmethod] + fn from_tiktoken(py: Python<'_>, encoding: &str) -> PyResult { + #[cfg(feature = "tiktoken")] + { + Self::load(py, || CoreTokenCounter::from_tiktoken(encoding)) + } + #[cfg(not(feature = "tiktoken"))] + { + let _ = (py, encoding); + Err(RustBridgeDeclined::new_err( + "tokenizer backend requires the tiktoken feature", + )) + } } fn acount_request<'py>(&self, py: Python<'py>, body: &[u8]) -> PyResult> { @@ -62,6 +116,7 @@ impl TokenCounter { } impl TokenCounter { + #[cfg(any(feature = "fast", feature = "huggingface", feature = "tiktoken"))] fn load( py: Python<'_>, load: impl FnOnce() -> Result + Send, @@ -74,6 +129,7 @@ impl TokenCounter { } } +#[cfg(any(feature = "fast", feature = "huggingface", feature = "tiktoken"))] fn encode_parallelism() -> usize { available_parallelism().map_or(1, NonZero::get) } @@ -86,7 +142,10 @@ fn count_body(counter: &CoreTokenCounter, body: &[u8]) -> Result PyErr { let message = error.to_string(); match error { - Error::Load(_) | Error::Ranks(_) | Error::UnicodeClasses => PyValueError::new_err(message), + Error::Load(_) + | Error::Ranks(_) + | Error::UnicodeClasses + | Error::UnsupportedTokenizer(_) => PyValueError::new_err(message), Error::RequestParse(_) | Error::MissingInput | Error::FloatText diff --git a/litellm-rust/crates/secrets-aws/Cargo.toml b/litellm-rust/crates/secrets-aws/Cargo.toml new file mode 100644 index 00000000000..b1dc5b33cda --- /dev/null +++ b/litellm-rust/crates/secrets-aws/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "litellm-secrets-aws" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-auth-aws.workspace = true +litellm-secrets-types.workspace = true +litellm-core-utils.workspace = true +serde_json.workspace = true +thiserror.workspace = true +tracing = "0.1" +veil.workspace = true +aws-sdk-kms = "1.120.0" +aws-sdk-secretsmanager = "1.117.0" +aws-credential-types = "1.3.0" + +[dev-dependencies] +base64.workspace = true +rstest.workspace = true +tokio.workspace = true +wiremock = "0.6.5" diff --git a/litellm-rust/crates/secrets-aws/src/auth.rs b/litellm-rust/crates/secrets-aws/src/auth.rs new file mode 100644 index 00000000000..954cfa2f8fd --- /dev/null +++ b/litellm-rust/crates/secrets-aws/src/auth.rs @@ -0,0 +1,79 @@ +use std::sync::Arc; + +use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future}; +use litellm_auth_aws::{ + AwsAuthConfig, + constants::{AWS_DEFAULT_REGION, AWS_REGION, AWS_REGION_NAME}, + resolve_credentials, +}; +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::KeyManagementSettings; + +use crate::Error; + +#[derive(Clone)] +pub(crate) struct Credentials { + config: AwsAuthConfig, + environment: Arc, +} + +impl Credentials { + pub(crate) fn new( + settings: &KeyManagementSettings, + environment: Arc, + ) -> Self { + Self { + config: AwsAuthConfig { + region_name: region(settings, environment.as_ref()).ok(), + role_name: settings.aws_role_name.clone(), + session_name: settings.aws_session_name.clone(), + external_id: settings + .aws_external_id + .as_ref() + .map(|v| v.expose().to_owned()), + profile_name: settings.aws_profile_name.clone(), + web_identity_token: settings + .aws_web_identity_token + .as_ref() + .map(|v| v.expose().to_owned()), + sts_endpoint: settings.aws_sts_endpoint.clone(), + ..Default::default() + }, + environment, + } + } +} + +impl ProvideCredentials for Credentials { + fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a> + where + Self: 'a, + { + future::ProvideCredentials::new(async { + resolve_credentials(self.config.clone(), &|name| self.environment.get(name)) + .await + .map_err(|_| { + CredentialsError::provider_error("secret manager authentication failed") + }) + }) + } +} + +pub(crate) fn region( + settings: &KeyManagementSettings, + environment: &dyn Lookup, +) -> Result { + settings + .aws_region_name + .clone() + .or_else(|| environment.get(AWS_REGION_NAME)) + .or_else(|| environment.get(AWS_REGION)) + .or_else(|| environment.get(AWS_DEFAULT_REGION)) + .ok_or(Error::MissingRegion) +} + +impl std::fmt::Debug for Credentials { + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + f.debug_struct("Credentials").finish_non_exhaustive() + } +} diff --git a/litellm-rust/crates/secrets-aws/src/error.rs b/litellm-rust/crates/secrets-aws/src/error.rs new file mode 100644 index 00000000000..23595397a13 --- /dev/null +++ b/litellm-rust/crates/secrets-aws/src/error.rs @@ -0,0 +1,31 @@ +use aws_sdk_secretsmanager::error::SdkError; + +#[derive(thiserror::Error, veil::Redact)] +pub enum Error { + #[error("AWS authentication failed")] + Auth(#[from] #[redact] litellm_auth_aws::Error), + #[error("AWS region is not configured")] + MissingRegion, + #[error("KMS response has no plaintext")] + MissingPlaintext, + #[error("AWS request timed out")] + Timeout, + #[error("AWS KMS decrypt failed")] + Decrypt(#[from] #[redact] Box>), + #[error("AWS Secrets Manager read failed")] + Read(#[from] #[redact] Box>), + #[error("AWS Secrets Manager create failed")] + Create(#[from] #[redact] Box>), + #[error("AWS Secrets Manager update failed")] + Put(#[from] #[redact] Box>), + #[error("AWS Secrets Manager delete failed")] + Delete(#[from] #[redact] Box>), + #[error("AWS Secrets Manager replication failed")] + Replicate(#[from] #[redact] Box>), + #[error("AWS Secrets Manager response has no string payload")] + MissingString, + #[error("primary secret is not a JSON object")] + PrimarySecret, + #[error(transparent)] + Operation(#[from] litellm_secrets_types::Error), +} diff --git a/litellm-rust/crates/secrets-aws/src/kms.rs b/litellm-rust/crates/secrets-aws/src/kms.rs new file mode 100644 index 00000000000..a66b1c4d2fe --- /dev/null +++ b/litellm-rust/crates/secrets-aws/src/kms.rs @@ -0,0 +1,63 @@ +use litellm_auth_aws::constants::AWS_REGION_NAME; +use std::sync::Arc; + +use aws_sdk_kms::{ + Client, + config::{BehaviorVersion, Region}, + primitives::Blob, +}; +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::KeyManagementSettings; + +use crate::{Error, auth}; + +#[derive(Clone)] +pub struct AwsKms { + client: Client, +} + +impl AwsKms { + pub fn new(client: Client) -> Self { + Self { client } + } + + pub async fn decrypt(&self, ciphertext: Vec) -> Result, Error> { + let response = self + .client + .decrypt() + .ciphertext_blob(Blob::new(ciphertext)) + .send() + .await + .map_err(|error| Error::Decrypt(Box::new(error)))?; + Ok(response + .plaintext + .ok_or(Error::MissingPlaintext)? + .into_inner()) + } +} + +pub fn validate_environment(environment: &dyn Lookup) -> Result<(), Error> { + environment + .get(AWS_REGION_NAME) + .map(|_| ()) + .ok_or(Error::MissingRegion) +} + +pub fn load_aws_kms( + use_aws_kms: Option, + settings: &KeyManagementSettings, + environment: Arc, +) -> Result, Error> { + if use_aws_kms != Some(true) { + return Ok(None); + } + if settings.aws_region_name.is_none() { + validate_environment(environment.as_ref())?; + } + let config = aws_sdk_kms::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new(auth::region(settings, environment.as_ref())?)) + .credentials_provider(auth::Credentials::new(settings, environment)) + .build(); + Ok(Some(AwsKms::new(Client::from_conf(config)))) +} diff --git a/litellm-rust/crates/secrets-aws/src/lib.rs b/litellm-rust/crates/secrets-aws/src/lib.rs new file mode 100644 index 00000000000..d17eb38c7eb --- /dev/null +++ b/litellm-rust/crates/secrets-aws/src/lib.rs @@ -0,0 +1,10 @@ +#![forbid(unsafe_code)] + +mod auth; +mod error; +pub mod kms; +pub mod secret_manager; + +pub use error::Error; +pub use kms::{AwsKms, load_aws_kms}; +pub use secret_manager::{AwsSecretWriteSettings, AwsSecretsManagerV2, RotationResponse}; diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager.rs b/litellm-rust/crates/secrets-aws/src/secret_manager.rs new file mode 100644 index 00000000000..493cb1d2e8f --- /dev/null +++ b/litellm-rust/crates/secrets-aws/src/secret_manager.rs @@ -0,0 +1,287 @@ +use litellm_auth_aws::constants::AWS_BEDROCK_RUNTIME_ENDPOINT; +use std::{collections::BTreeMap, sync::Arc}; + +use aws_sdk_secretsmanager::{ + Client, + config::{BehaviorVersion, Region}, + operation::{ + create_secret::CreateSecretOutput, delete_secret::DeleteSecretOutput, + put_secret_value::PutSecretValueOutput, + replicate_secret_to_regions::ReplicateSecretToRegionsOutput, + }, + types::{ReplicaRegionType, Tag}, +}; +use litellm_auth_aws::constants::{ + AWS_ACCESS_KEY_ID, AWS_REGION, AWS_REGION_NAME, AWS_SECRET_ACCESS_KEY, +}; +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::{ + BaseSecretManager, KeyManagementSettings, Secret, SecretValue, async_rotate_secret, +}; +use serde_json::Value; + +use crate::{Error, auth}; + +#[derive(Clone)] +pub struct AwsSecretsManagerV2 { + client: Client, + write_settings: AwsSecretWriteSettings, +} + +#[derive(Clone, Debug, Default)] +pub struct AwsSecretWriteSettings { + pub kms_key_id: Option, + pub tags: Option>, + pub replica_regions: Option>, +} + +impl From<&KeyManagementSettings> for AwsSecretWriteSettings { + fn from(settings: &KeyManagementSettings) -> Self { + Self { + kms_key_id: settings.kms_key_id.clone(), + tags: settings.tags.clone(), + replica_regions: settings.replica_regions.clone(), + } + } +} + +#[derive(Debug)] +pub enum RotationResponse { + Created(CreateSecretOutput), + Updated(PutSecretValueOutput), +} + +impl AwsSecretsManagerV2 { + pub fn new(client: Client, write_settings: AwsSecretWriteSettings) -> Self { + Self { + client, + write_settings, + } + } + + pub fn load_aws_secret_manager( + use_aws_secret_manager: Option, + settings: KeyManagementSettings, + environment: Arc, + ) -> Result, Error> { + if use_aws_secret_manager != Some(true) { + return Ok(None); + } + let builder = aws_sdk_secretsmanager::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new(auth::region(&settings, environment.as_ref())?)) + .credentials_provider(auth::Credentials::new(&settings, environment.clone())); + let config = match environment.get(AWS_BEDROCK_RUNTIME_ENDPOINT) { + Some(url) => builder + .endpoint_url(url.replace("bedrock-runtime", "secretsmanager")) + .build(), + None => builder.build(), + }; + Ok(Some(Self::new( + Client::from_conf(config), + (&settings).into(), + ))) + } + + pub async fn read_secret_for_resolver( + &self, + name: &str, + primary_name: Option<&str>, + environment: &(dyn Lookup + Sync), + ) -> Result, Error> { + if bootstrap_key(name) { + return Ok(environment + .get(name) + .map(SecretValue::new) + .map(Secret::String)); + } + match primary_name.filter(|name| !name.is_empty()) { + None => self + .async_read_secret(name) + .await + .map(|value| value.map(Secret::String)), + Some(primary) => { + let value = if bootstrap_key(primary) { + environment.get(primary).map(SecretValue::new) + } else { + self.async_read_secret(primary).await? + }; + let Some(value) = value else { + return Ok(None); + }; + let object: Value = + serde_json::from_str(value.expose()).map_err(|_| Error::PrimarySecret)?; + let object = object.as_object().ok_or(Error::PrimarySecret)?; + Ok(object.get(name).cloned().map(Secret::from_json)) + } + } + } + + pub async fn async_read_secret(&self, name: &str) -> Result, Error> { + match self.client.get_secret_value().secret_id(name).send().await { + Ok(response) => response + .secret_string + .map(SecretValue::new) + .map(Some) + .ok_or(Error::MissingString), + Err(error) + if matches!( + &error, + aws_sdk_secretsmanager::error::SdkError::TimeoutError(_) + ) || matches!(&error, aws_sdk_secretsmanager::error::SdkError::DispatchFailure(failure) if failure.is_timeout()) => + { + Err(Error::Timeout) + } + Err(error) + if error + .as_service_error() + .is_some_and(|error| error.is_resource_not_found_exception()) => + { + Ok(None) + } + Err(error) => Err(Error::Read(Box::new(error))), + } + } + + pub async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + description: Option<&str>, + ) -> Result { + let response = self + .client + .create_secret() + .name(name) + .secret_string(value.expose()) + .set_description(description.filter(|v| !v.is_empty()).map(str::to_owned)) + .set_kms_key_id( + self.write_settings + .kms_key_id + .clone() + .filter(|v| !v.is_empty()), + ) + .set_tags(self.write_settings.tags.as_ref().map(|tags| { + tags.iter() + .map(|(key, value)| Tag::builder().key(key).value(value).build()) + .collect() + })) + .send() + .await + .map_err(|error| Error::Create(Box::new(error)))?; + if let Some(regions) = &self.write_settings.replica_regions + && !regions.is_empty() + && self.async_replicate_secret(name, regions).await.is_err() + { + tracing::warn!("secret created but replication failed"); + } + Ok(response) + } + + pub async fn async_replicate_secret( + &self, + name: &str, + regions: &[String], + ) -> Result, Error> { + if regions.is_empty() { + return Ok(None); + } + self.client + .replicate_secret_to_regions() + .secret_id(name) + .set_add_replica_regions(Some( + regions + .iter() + .map(|region| ReplicaRegionType::builder().region(region).build()) + .collect(), + )) + .send() + .await + .map(Some) + .map_err(|error| Error::Replicate(Box::new(error))) + } + + pub async fn async_put_secret_value( + &self, + name: &str, + value: &SecretValue, + ) -> Result { + self.client + .put_secret_value() + .secret_id(name) + .secret_string(value.expose()) + .send() + .await + .map_err(|error| Error::Put(Box::new(error))) + } + + pub async fn async_delete_secret( + &self, + name: &str, + recovery_window_in_days: i64, + ) -> Result { + self.client + .delete_secret() + .secret_id(name) + .recovery_window_in_days(recovery_window_in_days) + .send() + .await + .map_err(|error| Error::Delete(Box::new(error))) + } + + pub async fn async_rotate_secret( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + ) -> Result { + if current_name == new_name { + return self + .async_put_secret_value(current_name, value) + .await + .map(RotationResponse::Updated); + } + async_rotate_secret(self, current_name, new_name, value) + .await + .map(RotationResponse::Created) + } +} + +impl BaseSecretManager for AwsSecretsManagerV2 { + type Error = Error; + type WriteResponse = CreateSecretOutput; + type DeleteResponse = DeleteSecretOutput; + + async fn async_read_secret(&self, name: &str) -> Result, Error> { + self.async_read_secret(name).await + } + + async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + description: Option<&str>, + ) -> Result { + self.async_write_secret(name, value, description).await + } + + async fn async_delete_secret( + &self, + name: &str, + recovery_window_in_days: i64, + ) -> Result { + self.async_delete_secret(name, recovery_window_in_days) + .await + } +} + +fn bootstrap_key(name: &str) -> bool { + matches!( + name, + AWS_ACCESS_KEY_ID + | AWS_SECRET_ACCESS_KEY + | AWS_REGION_NAME + | AWS_REGION + | AWS_BEDROCK_RUNTIME_ENDPOINT + ) +} diff --git a/litellm-rust/crates/secrets-aws/tests/kms.rs b/litellm-rust/crates/secrets-aws/tests/kms.rs new file mode 100644 index 00000000000..39a50297551 --- /dev/null +++ b/litellm-rust/crates/secrets-aws/tests/kms.rs @@ -0,0 +1,59 @@ +use aws_sdk_kms::{ + Client, + config::{BehaviorVersion, Credentials, Region, retry::RetryConfig}, +}; +use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_secrets_aws::AwsKms; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_json, header}, +}; + +#[tokio::test] +async fn kms_decrypt_calls_the_sdk_without_applying_lookup_policy() { + let server = MockServer::start().await; + let plaintext = " private-value\n"; + Mock::given(header("x-amz-target", "TrentService.Decrypt")) + .and(body_json( + serde_json::json!({"CiphertextBlob": STANDARD.encode("encrypted")}), + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"Plaintext": STANDARD.encode(plaintext)})), + ) + .expect(1) + .mount(&server) + .await; + let client = Client::from_conf( + aws_sdk_kms::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new("us-east-1")) + .credentials_provider(Credentials::new("test", "test", None, None, "test")) + .endpoint_url(server.uri()) + .retry_config(RetryConfig::disabled()) + .build(), + ); + let manager = AwsKms::new(client); + assert_eq!( + manager.decrypt(b"encrypted".to_vec()).await.unwrap(), + plaintext.as_bytes() + ); +} + +#[test] +fn disabled_kms_loader_does_not_require_environment_configuration() { + use litellm_secrets_aws::load_aws_kms; + use litellm_secrets_types::KeyManagementSettings; + use std::sync::Arc; + for enabled in [None, Some(false)] { + assert!( + load_aws_kms( + enabled, + &KeyManagementSettings::default(), + Arc::new(|_: &str| None) + ) + .unwrap() + .is_none() + ); + } +} diff --git a/litellm-rust/crates/secrets-aws/tests/secret_manager.rs b/litellm-rust/crates/secrets-aws/tests/secret_manager.rs new file mode 100644 index 00000000000..a410767cb5a --- /dev/null +++ b/litellm-rust/crates/secrets-aws/tests/secret_manager.rs @@ -0,0 +1,312 @@ +use std::sync::atomic::{AtomicUsize, Ordering}; + +use aws_sdk_secretsmanager::{ + Client, + config::{BehaviorVersion, Credentials, Region, retry::RetryConfig}, +}; +use litellm_secrets_aws::{AwsSecretsManagerV2, Error, RotationResponse}; +use litellm_secrets_types::{KeyManagementSettings, SecretValue}; +use serde_json::json; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_partial_json, header}, +}; + +fn manager(server: &MockServer, settings: KeyManagementSettings) -> AwsSecretsManagerV2 { + let client = Client::from_conf( + aws_sdk_secretsmanager::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new("us-east-1")) + .credentials_provider(Credentials::new("test", "test", None, None, "test")) + .endpoint_url(server.uri()) + .retry_config(RetryConfig::disabled()) + .build(), + ); + AwsSecretsManagerV2::new(client, (&settings).into()) +} + +#[rstest::rstest] +#[case::string_value("KEY", Some("value"))] +#[case::missing_value("missing", None)] +#[case::non_string_value("BOOL", None)] +#[tokio::test] +async fn primary_lookup_preserves_read_semantics( + #[case] name: &str, + #[case] expected: Option<&str>, +) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .and(body_partial_json(json!({"SecretId":"primary"}))) + .respond_with( + ResponseTemplate::new(200).set_body_json( + json!({"SecretString":json!({"KEY":"value", "BOOL":true}).to_string()}), + ), + ) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, KeyManagementSettings::default()); + assert_eq!( + manager + .read_secret_for_resolver(name, Some("primary"), &|_: &str| None) + .await + .unwrap() + .and_then(|v| v.as_str().map(str::to_owned)) + .as_deref(), + expected + ); +} + +#[rstest::rstest] +#[case::access_key("AWS_ACCESS_KEY_ID")] +#[case::secret_access_key("AWS_SECRET_ACCESS_KEY")] +#[case::region_name("AWS_REGION_NAME")] +#[case::region("AWS_REGION")] +#[case::bedrock_endpoint("AWS_BEDROCK_RUNTIME_ENDPOINT")] +#[tokio::test] +async fn bootstrap_keys_bypass_primary_lookup(#[case] name: &str) { + let server = MockServer::start().await; + let manager = manager(&server, KeyManagementSettings::default()); + assert_eq!( + manager + .read_secret_for_resolver(name, Some("primary"), &|_: &str| Some("bootstrap".into())) + .await + .unwrap() + .unwrap() + .as_str() + .unwrap(), + "bootstrap" + ); +} + +#[tokio::test] +async fn failed_read_returns_none_but_invalid_primary_json_is_an_error() { + let server = MockServer::start().await; + Mock::given(body_partial_json(json!({"SecretId":"missing"}))) + .respond_with( + ResponseTemplate::new(400).set_body_json(json!({"__type":"ResourceNotFoundException"})), + ) + .mount(&server) + .await; + Mock::given(body_partial_json(json!({"SecretId":"invalid"}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"not-json"}))) + .mount(&server) + .await; + let manager = manager(&server, KeyManagementSettings::default()); + assert!( + manager + .async_read_secret("missing") + .await + .unwrap() + .is_none() + ); + assert!(matches!( + manager + .read_secret_for_resolver("KEY", Some("invalid"), &|_: &str| None) + .await, + Err(Error::PrimarySecret) + )); +} + +#[tokio::test] +async fn same_name_rotation_uses_put_and_returns_its_response() { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.PutSecretValue")) + .and(body_partial_json( + json!({"SecretId":"key", "SecretString":"replacement"}), + )) + .respond_with( + ResponseTemplate::new(200).set_body_json(json!({"Name":"key", "VersionId":"version"})), + ) + .expect(1) + .mount(&server) + .await; + let response = manager(&server, KeyManagementSettings::default()) + .async_rotate_secret("key", "key", &SecretValue::new("replacement")) + .await + .unwrap(); + match response { + RotationResponse::Updated(output) => assert_eq!(output.version_id(), Some("version")), + _ => panic!("rotation created a second secret"), + } + assert_eq!(server.received_requests().await.unwrap().len(), 1); +} + +#[tokio::test] +async fn renamed_rotation_reads_creates_verifies_then_deletes() { + let server = MockServer::start().await; + let step = AtomicUsize::new(0); + Mock::given(wiremock::matchers::method("POST")) + .respond_with(move |request: &wiremock::Request| { + let body: serde_json::Value = request.body_json().unwrap(); + let action = request + .headers + .get("x-amz-target") + .unwrap() + .to_str() + .unwrap(); + match step.fetch_add(1, Ordering::SeqCst) { + 0 => { + assert_eq!(action, "secretsmanager.GetSecretValue"); + assert_eq!(body["SecretId"], "old"); + ResponseTemplate::new(200).set_body_json(json!({"SecretString":"old-value"})) + } + 1 => { + assert_eq!(action, "secretsmanager.CreateSecret"); + assert_eq!(body["Name"], "new"); + assert_eq!(body["Description"], "Rotated from old"); + assert_eq!(body["SecretString"], "replacement"); + ResponseTemplate::new(200).set_body_json(json!({"Name":"new"})) + } + 2 => { + assert_eq!(action, "secretsmanager.GetSecretValue"); + assert_eq!(body["SecretId"], "new"); + ResponseTemplate::new(200).set_body_json(json!({"SecretString":"replacement"})) + } + 3 => { + assert_eq!(action, "secretsmanager.DeleteSecret"); + assert_eq!(body["SecretId"], "old"); + assert_eq!(body["RecoveryWindowInDays"], 7); + ResponseTemplate::new(200).set_body_json(json!({"Name":"old"})) + } + _ => panic!("unexpected request"), + } + }) + .expect(4) + .mount(&server) + .await; + assert!(matches!( + manager(&server, KeyManagementSettings::default()) + .async_rotate_secret("old", "new", &SecretValue::new("replacement")) + .await + .unwrap(), + RotationResponse::Created(_) + )); +} + +#[tokio::test] +async fn creation_passes_tags_and_kms_and_survives_replication_failure() { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.CreateSecret")) + .and(body_partial_json(json!({"Name":"key", "SecretString":"value", "KmsKeyId":"kms-key", "Tags":[{"Key":"stage", "Value":"test"}]}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"key"}))).expect(1).mount(&server).await; + Mock::given(header( + "x-amz-target", + "secretsmanager.ReplicateSecretToRegions", + )) + .and(body_partial_json( + json!({"SecretId":"key", "AddReplicaRegions":[{"Region":"replica-region"}]}), + )) + .respond_with( + ResponseTemplate::new(400).set_body_json(json!({"__type":"InvalidRequestException"})), + ) + .expect(1) + .mount(&server) + .await; + let settings = KeyManagementSettings { + kms_key_id: Some("kms-key".into()), + tags: Some(std::collections::BTreeMap::from([( + "stage".into(), + "test".into(), + )])), + replica_regions: Some(vec!["replica-region".into()]), + ..Default::default() + }; + let manager = manager(&server, settings); + assert_eq!( + manager + .async_write_secret("key", &SecretValue::new("value"), None) + .await + .unwrap() + .name(), + Some("key") + ); + assert!( + manager + .async_replicate_secret("key", &[]) + .await + .unwrap() + .is_none() + ); +} + +#[tokio::test] +async fn credential_failures_are_not_swallowed_as_missing_secrets() { + use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future}; + #[derive(Debug)] + struct FailedCredentials; + impl ProvideCredentials for FailedCredentials { + fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a> + where + Self: 'a, + { + future::ProvideCredentials::ready(Err(CredentialsError::provider_error( + "private-auth-detail", + ))) + } + } + let server = MockServer::start().await; + let config = aws_sdk_secretsmanager::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new("us-east-1")) + .credentials_provider(FailedCredentials) + .endpoint_url(server.uri()) + .retry_config(RetryConfig::disabled()) + .build(); + let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default()); + let error = manager.async_read_secret("key").await.unwrap_err(); + assert!(!format!("{error:?}").contains("private-auth-detail")); + assert!(matches!(error, Error::Read(_))); + assert!(server.received_requests().await.unwrap().is_empty()); +} + +#[tokio::test] +async fn read_timeout_is_an_error_and_cannot_be_mistaken_for_missing() { + use std::time::Duration; + let server = MockServer::start().await; + Mock::given(wiremock::matchers::method("POST")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_secs(1)) + .set_body_json(json!({"SecretString":"late"})), + ) + .mount(&server) + .await; + let config = aws_sdk_secretsmanager::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new("us-east-1")) + .credentials_provider(Credentials::new("test", "test", None, None, "test")) + .endpoint_url(server.uri()) + .retry_config(RetryConfig::disabled()) + .timeout_config( + aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder() + .operation_timeout(Duration::from_millis(30)) + .build(), + ) + .build(); + let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default()); + assert!(matches!( + manager.async_read_secret("key").await, + Err(Error::Timeout) + )); +} + +#[rstest::rstest] +#[case::denied(400, "AccessDeniedException")] +#[case::throttled(400, "ThrottlingException")] +#[case::unavailable(503, "ServiceUnavailableException")] +#[tokio::test] +async fn service_failures_remain_errors(#[case] status: u16, #[case] code: &str) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .respond_with(ResponseTemplate::new(status).set_body_json(json!({"__type":code}))) + .expect(1) + .mount(&server) + .await; + assert!(matches!( + manager(&server, KeyManagementSettings::default()) + .async_read_secret("key") + .await, + Err(Error::Read(_)) + )); +} diff --git a/litellm-rust/crates/secrets-google/Cargo.toml b/litellm-rust/crates/secrets-google/Cargo.toml new file mode 100644 index 00000000000..daecf20ff9e --- /dev/null +++ b/litellm-rust/crates/secrets-google/Cargo.toml @@ -0,0 +1,28 @@ +[package] +name = "litellm-secrets-google" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-auth-gcp = { workspace = true, features = ["google-sdk"] } +litellm-secrets-types.workspace = true +litellm-auth-types.workspace = true +litellm-core-utils.workspace = true +base64.workspace = true +serde_json.workspace = true +thiserror.workspace = true +moka.workspace = true +veil.workspace = true +google-cloud-kms-v1 = "1.14.0" +google-cloud-gax = { version = "1.14.0", default-features = false } +percent-encoding = "2.3" +serde.workspace = true +reqwest.workspace = true + +[dev-dependencies] +google-cloud-auth.workspace = true +rstest.workspace = true +tokio.workspace = true +wiremock = "0.6.5" diff --git a/litellm-rust/crates/secrets-google/src/auth.rs b/litellm-rust/crates/secrets-google/src/auth.rs new file mode 100644 index 00000000000..45fc99d8d5f --- /dev/null +++ b/litellm-rust/crates/secrets-google/src/auth.rs @@ -0,0 +1,21 @@ +use std::sync::Arc; + +use litellm_auth_gcp::{GoogleCredentials, VertexConfig}; +use litellm_auth_types::{InputSource, Sourced}; +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::SecretValue; + +pub(crate) fn credentials( + project: Option, + credentials: Option, + environment: Arc, +) -> GoogleCredentials { + GoogleCredentials::new( + VertexConfig::new( + credentials.map(|value| Sourced::new(value, InputSource::Environment)), + project, + None, + ), + Arc::new(move |name| environment.get(name)), + ) +} diff --git a/litellm-rust/crates/secrets-google/src/error.rs b/litellm-rust/crates/secrets-google/src/error.rs new file mode 100644 index 00000000000..a94cc8c2de6 --- /dev/null +++ b/litellm-rust/crates/secrets-google/src/error.rs @@ -0,0 +1,43 @@ +#[derive(thiserror::Error, veil::Redact)] +pub enum Error { + #[error("Google KMS client configuration failed")] + Client( + #[from] + #[redact] + google_cloud_gax::client_builder::Error, + ), + #[error("Google authentication failed")] + Auth( + #[from] + #[redact] + litellm_auth_types::Error, + ), + #[error("Google KMS request failed")] + Kms( + #[from] + #[redact] + google_cloud_gax::error::Error, + ), + #[error("Google Secret Manager HTTP request failed")] + Http( + #[from] + #[redact] + reqwest::Error, + ), + #[error("Google Secret Manager returned HTTP {0}")] + Status(u16), + #[error("Google Secret Manager returned no payload")] + MissingPayload, + #[error("required environment variable is missing: {0}")] + MissingEnvironment(&'static str), + #[error("invalid refresh interval")] + RefreshInterval, + #[error("payload is not valid base64")] + Base64(#[from] base64::DecodeError), + #[error("decrypted value is not UTF-8")] + Utf8, + #[error("invalid Google Secret Manager endpoint")] + Endpoint, + #[error("Google Secret Manager requires an enterprise license")] + EnterpriseRequired, +} diff --git a/litellm-rust/crates/secrets-google/src/kms.rs b/litellm-rust/crates/secrets-google/src/kms.rs new file mode 100644 index 00000000000..3a247edaa35 --- /dev/null +++ b/litellm-rust/crates/secrets-google/src/kms.rs @@ -0,0 +1,67 @@ +use std::sync::Arc; + +use google_cloud_kms_v1::client::KeyManagementService; +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::SecretValue; + +use crate::{Error, auth}; + +const GOOGLE_APPLICATION_CREDENTIALS: &str = "GOOGLE_APPLICATION_CREDENTIALS"; +const GOOGLE_KMS_RESOURCE_NAME: &str = "GOOGLE_KMS_RESOURCE_NAME"; + +#[derive(Clone)] +pub struct GoogleKms { + client: KeyManagementService, + resource_name: String, +} + +impl GoogleKms { + pub fn new(client: KeyManagementService, resource_name: String) -> Self { + Self { + client, + resource_name, + } + } + + pub async fn decrypt(&self, ciphertext: Vec) -> Result, Error> { + let response = self + .client + .decrypt() + .set_name(&self.resource_name) + .set_ciphertext(ciphertext) + .send() + .await?; + Ok(response.plaintext.to_vec()) + } +} + +pub fn validate_environment(environment: &dyn Lookup) -> Result<(), Error> { + for key in [GOOGLE_APPLICATION_CREDENTIALS, GOOGLE_KMS_RESOURCE_NAME] { + if environment.get(key).is_none() { + return Err(Error::MissingEnvironment(key)); + } + } + Ok(()) +} + +pub async fn load_google_kms( + use_google_kms: Option, + environment: Arc, +) -> Result, Error> { + if use_google_kms != Some(true) { + return Ok(None); + } + validate_environment(environment.as_ref())?; + let credentials = environment + .get(GOOGLE_APPLICATION_CREDENTIALS) + .ok_or(Error::MissingEnvironment(GOOGLE_APPLICATION_CREDENTIALS))?; + let resource_name = environment + .get(GOOGLE_KMS_RESOURCE_NAME) + .ok_or(Error::MissingEnvironment(GOOGLE_KMS_RESOURCE_NAME))?; + let credentials = auth::credentials(None, Some(SecretValue::new(credentials)), environment); + let client = KeyManagementService::builder() + .with_credentials(credentials) + .build() + .await?; + Ok(Some(GoogleKms::new(client, resource_name))) +} diff --git a/litellm-rust/crates/secrets-google/src/lib.rs b/litellm-rust/crates/secrets-google/src/lib.rs new file mode 100644 index 00000000000..a664f11a018 --- /dev/null +++ b/litellm-rust/crates/secrets-google/src/lib.rs @@ -0,0 +1,10 @@ +#![forbid(unsafe_code)] + +mod auth; +mod error; +pub mod kms; +pub mod secret_manager; + +pub use error::Error; +pub use kms::{GoogleKms, load_google_kms}; +pub use secret_manager::GoogleSecretManager; diff --git a/litellm-rust/crates/secrets-google/src/secret_manager.rs b/litellm-rust/crates/secrets-google/src/secret_manager.rs new file mode 100644 index 00000000000..3c34d9cbcc4 --- /dev/null +++ b/litellm-rust/crates/secrets-google/src/secret_manager.rs @@ -0,0 +1,157 @@ +use std::{sync::Arc, time::Duration}; + +use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_core_utils::settings::Lookup; +use litellm_secrets_types::{Secret, SecretValue}; +use moka::future::Cache; +use serde::Deserialize; + +use litellm_auth_gcp::GoogleCredentials; + +use crate::{Error, auth}; + +const GOOGLE_SECRET_MANAGER_PROJECT_ID: &str = "GOOGLE_SECRET_MANAGER_PROJECT_ID"; +const GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL: &str = "GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL"; +const SECRET_MANAGER_REFRESH_INTERVAL: &str = "SECRET_MANAGER_REFRESH_INTERVAL"; +const GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER: &str = + "GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER"; +const GCS_PATH_SERVICE_ACCOUNT: &str = "GCS_PATH_SERVICE_ACCOUNT"; +const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(86400); +const DEFAULT_CACHE_TTL: Duration = Duration::from_secs(600); +const CACHE_CAPACITY: u64 = 200; + +#[derive(Clone)] +pub struct GoogleSecretManager { + client: reqwest::Client, + credentials: Arc, + endpoint: reqwest::Url, + project: String, + cache: Cache, + always_read: bool, +} + +#[derive(Deserialize)] +struct Response { + payload: Option, +} + +#[derive(Deserialize)] +struct Payload { + data: Option, +} + +impl GoogleSecretManager { + pub fn with_client( + client: reqwest::Client, + endpoint: reqwest::Url, + project: String, + environment: Arc, + refresh_interval: Option, + always_read: bool, + ) -> Result { + let credentials = auth::credentials( + Some(project.clone()), + environment + .get(GCS_PATH_SERVICE_ACCOUNT) + .map(SecretValue::new), + environment, + ); + let ttl = refresh_interval + .filter(|ttl| !ttl.is_zero()) + .unwrap_or(DEFAULT_CACHE_TTL); + let cache = Cache::builder() + .max_capacity(CACHE_CAPACITY) + .time_to_live(ttl) + .build(); + Ok(Self { + client, + credentials: Arc::new(credentials), + endpoint, + project, + cache, + always_read, + }) + } + + pub fn new( + environment: Arc, + enterprise_enabled: bool, + ) -> Result { + if !enterprise_enabled { + return Err(Error::EnterpriseRequired); + } + let project = environment + .get(GOOGLE_SECRET_MANAGER_PROJECT_ID) + .ok_or(Error::MissingEnvironment(GOOGLE_SECRET_MANAGER_PROJECT_ID))?; + let ttl = environment + .get(GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL) + .filter(|v| !v.is_empty()) + .map(|v| v.parse::().map_err(|_| Error::RefreshInterval)) + .transpose()? + .unwrap_or( + environment + .get(SECRET_MANAGER_REFRESH_INTERVAL) + .map(|v| v.parse::().map_err(|_| Error::RefreshInterval)) + .transpose()? + .unwrap_or(DEFAULT_REFRESH_INTERVAL.as_secs() as i64), + ); + let always_read = environment + .get(GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER) + .is_some_and(|v| v.eq_ignore_ascii_case("true")); + Self::with_client( + reqwest::Client::new(), + reqwest::Url::parse("https://secretmanager.googleapis.com").expect("static URL"), + project, + environment, + Some(if ttl < 0 { + Duration::from_nanos(1) + } else { + Duration::from_secs(ttl as u64) + }), + always_read, + ) + } + + pub async fn get_secret_from_google_secret_manager( + &self, + name: &str, + ) -> Result, Error> { + if !self.always_read + && let Some(cached) = self.cache.get(name).await + { + return Ok(Some(Secret::String(cached))); + } + let url = self + .endpoint + .join(&format!( + "/v1/projects/{}/secrets/{}/versions/latest:access", + percent_encoding::utf8_percent_encode( + &self.project, + percent_encoding::NON_ALPHANUMERIC + ), + percent_encoding::utf8_percent_encode(name, percent_encoding::NON_ALPHANUMERIC) + )) + .map_err(|_| Error::Endpoint)?; + let response = self + .client + .get(url) + .headers(self.credentials.request_headers().await?) + .send() + .await?; + if response.status() == reqwest::StatusCode::NOT_FOUND { + return Ok(None); + } + if response.status() != reqwest::StatusCode::OK { + return Err(Error::Status(response.status().as_u16())); + } + let response: Response = response.json().await?; + let Some(data) = response.payload.and_then(|payload| payload.data) else { + return Err(Error::MissingPayload); + }; + let bytes = STANDARD.decode(data)?; + let plaintext = String::from_utf8(bytes).map_err(|_| Error::Utf8)?; + let value = SecretValue::new(plaintext); + self.cache.insert(name.to_owned(), value.clone()).await; + Ok(Some(Secret::String(value))) + } +} diff --git a/litellm-rust/crates/secrets-google/tests/kms.rs b/litellm-rust/crates/secrets-google/tests/kms.rs new file mode 100644 index 00000000000..667ecd268c8 --- /dev/null +++ b/litellm-rust/crates/secrets-google/tests/kms.rs @@ -0,0 +1,49 @@ +use base64::{Engine, engine::general_purpose::STANDARD}; +use google_cloud_kms_v1::client::KeyManagementService; +use litellm_secrets_google::GoogleKms; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_json, path}, +}; + +#[tokio::test] +async fn google_kms_decrypts_using_the_configured_resource() { + let server = MockServer::start().await; + let resource = "projects/project/locations/global/keyRings/ring/cryptoKeys/key"; + Mock::given(path(format!("/v1/{resource}:decrypt"))) + .and(body_json( + serde_json::json!({"ciphertext":STANDARD.encode("encrypted")}), + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"plaintext":STANDARD.encode(" value\n")})), + ) + .expect(1) + .mount(&server) + .await; + let client = KeyManagementService::builder() + .with_endpoint(server.uri()) + .with_credentials(google_cloud_auth::credentials::anonymous::Builder::new().build()) + .with_retry_policy(google_cloud_gax::retry_policy::NeverRetry) + .build() + .await + .unwrap(); + let manager = GoogleKms::new(client, resource.into()); + assert_eq!( + manager.decrypt(b"encrypted".to_vec()).await.unwrap(), + b" value\n" + ); +} + +#[tokio::test] +async fn disabled_google_kms_loader_does_not_require_environment_configuration() { + use std::sync::Arc; + for enabled in [None, Some(false)] { + assert!( + litellm_secrets_google::load_google_kms(enabled, Arc::new(|_: &str| None)) + .await + .unwrap() + .is_none() + ); + } +} diff --git a/litellm-rust/crates/secrets-google/tests/secret_manager.rs b/litellm-rust/crates/secrets-google/tests/secret_manager.rs new file mode 100644 index 00000000000..b3b1d29e62c --- /dev/null +++ b/litellm-rust/crates/secrets-google/tests/secret_manager.rs @@ -0,0 +1,188 @@ +use std::{sync::Arc, time::Duration}; + +use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_secrets_google::{Error, GoogleSecretManager}; + +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{header, path}, +}; + +fn manager(server: &MockServer, always_read: bool, ttl: Duration) -> GoogleSecretManager { + GoogleSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + "project".into(), + Arc::new(|name: &str| (name == "VERTEX_AI_API_KEY").then(|| "token".into())), + Some(ttl), + always_read, + ) + .unwrap() +} + +#[rstest::rstest] +#[case::nonempty("private-value")] +#[case::empty("")] +#[tokio::test] +async fn successful_reads_use_auth_latest_version_and_cache_including_empty_values( + #[case] value: &str, +) { + let server = MockServer::start().await; + Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .and(header("authorization", "Bearer token")) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode(value)}})), + ) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, false, Duration::from_secs(60)); + for _ in 0..2 { + assert_eq!( + manager + .get_secret_from_google_secret_manager("key") + .await + .unwrap() + .unwrap() + .as_str() + .unwrap(), + value + ); + } +} + +#[rstest::rstest] +#[case::not_found(404, serde_json::json!({}))] +#[case::unauthorized(401, serde_json::json!({}))] +#[case::forbidden(403, serde_json::json!({}))] +#[case::throttled(429, serde_json::json!({}))] +#[case::unavailable(503, serde_json::json!({}))] +#[case::missing_payload(200, serde_json::json!({"payload":{}}))] +#[case::invalid_base64(200, serde_json::json!({"payload":{"data":"%%%"}}))] +#[tokio::test] +async fn failed_or_missing_reads_are_not_cached( + #[case] status: u16, + #[case] body: serde_json::Value, +) { + let server = MockServer::start().await; + let manager = manager(&server, false, Duration::from_secs(60)); + let failing = Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with(ResponseTemplate::new(status).set_body_json(body)) + .expect(1) + .mount_as_scoped(&server) + .await; + let result = manager.get_secret_from_google_secret_manager("key").await; + match status { + 404 => assert_eq!(result.unwrap(), None), + 200 => assert!(matches!( + result, + Err(Error::MissingPayload | Error::Base64(_)) + )), + status => assert!(matches!(result, Err(Error::Status(actual)) if actual == status)), + } + drop(failing); + Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode("recovered")}})), + ) + .expect(1) + .mount(&server) + .await; + for _ in 0..2 { + assert_eq!( + manager + .get_secret_from_google_secret_manager("key") + .await + .unwrap() + .unwrap() + .as_str(), + Some("recovered") + ); + } +} + +#[rstest::rstest] +#[case::always_read(true, Duration::from_secs(60))] +#[case::expired_cache(false, Duration::from_millis(1))] +#[tokio::test] +async fn always_read_and_expired_cache_fetch_again( + #[case] always_read: bool, + #[case] ttl: Duration, +) { + let server = MockServer::start().await; + Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode("value")}})), + ) + .expect(2) + .mount(&server) + .await; + let manager = manager(&server, always_read, ttl); + for _ in 0..2 { + tokio::time::sleep(Duration::from_millis(5)).await; + assert!( + manager + .get_secret_from_google_secret_manager("key") + .await + .unwrap() + .is_some() + ); + } +} + +#[test] +fn google_manager_requires_host_license_and_project_configuration() { + assert!(matches!( + GoogleSecretManager::new(Arc::new(|_: &str| None), false), + Err(Error::EnterpriseRequired) + )); + assert!(matches!( + GoogleSecretManager::new(Arc::new(|_: &str| None), true), + Err(Error::MissingEnvironment( + "GOOGLE_SECRET_MANAGER_PROJECT_ID" + )) + )); +} + +#[rstest::rstest] +#[case("true")] +#[case("null")] +#[case("\"text\"")] +#[case("{\"key\":1}")] +#[tokio::test] +async fn cache_preserves_raw_values(#[case] raw: &str) { + let server = MockServer::start().await; + Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode(raw)}})), + ) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, false, Duration::from_secs(60)); + for _ in 0..2 { + assert_eq!( + manager + .get_secret_from_google_secret_manager("key") + .await + .unwrap() + .unwrap() + .as_str(), + Some(raw) + ); + } +} diff --git a/litellm-rust/crates/secrets-types/Cargo.toml b/litellm-rust/crates/secrets-types/Cargo.toml new file mode 100644 index 00000000000..acd29746722 --- /dev/null +++ b/litellm-rust/crates/secrets-types/Cargo.toml @@ -0,0 +1,17 @@ +[package] +name = "litellm-secrets-types" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +litellm-auth-types.workspace = true +serde.workspace = true +serde_json.workspace = true +thiserror.workspace = true +veil.workspace = true + +[dev-dependencies] +rstest.workspace = true +tokio.workspace = true diff --git a/litellm-rust/crates/secrets-types/src/base_secret_manager.rs b/litellm-rust/crates/secrets-types/src/base_secret_manager.rs new file mode 100644 index 00000000000..d71bce64221 --- /dev/null +++ b/litellm-rust/crates/secrets-types/src/base_secret_manager.rs @@ -0,0 +1,58 @@ +use crate::{Error, SecretValue}; + +pub fn validate_secret_name(name: &str) -> Result<(), Error> { + if name.split('/').any(|segment| segment == "..") + || name + .chars() + .any(|c| c.is_control() || matches!(c, '\u{2028}' | '\u{2029}')) + { + return Err(Error::UnsafeSecretName); + } + Ok(()) +} + +#[expect( + async_fn_in_trait, + reason = "closed backend dispatch does not require Send bounds on generic rotation" +)] +pub trait BaseSecretManager { + type Error: From; + type WriteResponse; + type DeleteResponse; + + async fn async_read_secret(&self, name: &str) -> Result, Self::Error>; + async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + description: Option<&str>, + ) -> Result; + async fn async_delete_secret( + &self, + name: &str, + recovery_window_in_days: i64, + ) -> Result; +} + +pub async fn async_rotate_secret( + manager: &M, + current_name: &str, + new_name: &str, + value: &SecretValue, +) -> Result { + if manager.async_read_secret(current_name).await?.is_none() { + return Err(Error::CurrentSecretMissing.into()); + } + let response = manager + .async_write_secret( + new_name, + value, + Some(&format!("Rotated from {current_name}")), + ) + .await?; + if manager.async_read_secret(new_name).await?.is_none() { + return Err(Error::NewSecretMissing.into()); + } + manager.async_delete_secret(current_name, 7).await?; + Ok(response) +} diff --git a/litellm-rust/crates/secrets-types/src/config.rs b/litellm-rust/crates/secrets-types/src/config.rs new file mode 100644 index 00000000000..36d319311a3 --- /dev/null +++ b/litellm-rust/crates/secrets-types/src/config.rs @@ -0,0 +1,92 @@ +use std::collections::BTreeMap; + +use serde::{Deserialize, Serialize}; + +use crate::SecretValue; + +#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum KeyManagementSystem { + GoogleKms, + AzureKeyVault, + AwsSecretManager, + GoogleSecretManager, + HashicorpVault, + Cyberark, + Local, + AwsKms, + Custom, +} + +#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)] +#[serde(rename_all = "snake_case")] +pub enum AccessMode { + #[default] + ReadOnly, + WriteOnly, + ReadAndWrite, +} + +impl AccessMode { + pub fn readable(self) -> bool { + matches!(self, Self::ReadOnly | Self::ReadAndWrite) + } +} + +#[derive(Clone, Debug, Deserialize, Serialize, PartialEq, Eq)] +#[serde(default)] +pub struct KeyManagementSettings { + pub hosted_keys: Option>, + pub store_virtual_keys: Option, + pub prefix_for_stored_virtual_keys: String, + pub access_mode: AccessMode, + pub primary_secret_name: Option, + pub description: Option, + pub tags: Option>, + pub kms_key_id: Option, + pub custom_secret_manager: Option, + pub aws_region_name: Option, + pub aws_role_name: Option, + pub aws_session_name: Option, + #[serde(serialize_with = "serialize_secret")] + pub aws_external_id: Option, + pub aws_profile_name: Option, + #[serde(serialize_with = "serialize_secret")] + pub aws_web_identity_token: Option, + pub aws_sts_endpoint: Option, + pub replica_regions: Option>, +} + +impl Default for KeyManagementSettings { + fn default() -> Self { + Self { + hosted_keys: None, + store_virtual_keys: Some(false), + prefix_for_stored_virtual_keys: "litellm/".into(), + access_mode: AccessMode::ReadOnly, + primary_secret_name: None, + description: None, + tags: None, + kms_key_id: None, + custom_secret_manager: None, + aws_region_name: None, + aws_role_name: None, + aws_session_name: None, + aws_external_id: None, + aws_profile_name: None, + aws_web_identity_token: None, + aws_sts_endpoint: None, + replica_regions: None, + } + } +} + +fn serialize_secret( + value: &Option, + serializer: S, +) -> Result { + value + .as_ref() + .map(SecretValue::expose) + .serialize(serializer) +} diff --git a/litellm-rust/crates/secrets-types/src/error.rs b/litellm-rust/crates/secrets-types/src/error.rs new file mode 100644 index 00000000000..cae9c7f4c69 --- /dev/null +++ b/litellm-rust/crates/secrets-types/src/error.rs @@ -0,0 +1,9 @@ +#[derive(Debug, thiserror::Error, PartialEq, Eq)] +pub enum Error { + #[error("secret name contains an unsafe path segment or control character")] + UnsafeSecretName, + #[error("current secret was not found")] + CurrentSecretMissing, + #[error("new secret could not be verified")] + NewSecretMissing, +} diff --git a/litellm-rust/crates/secrets-types/src/lib.rs b/litellm-rust/crates/secrets-types/src/lib.rs new file mode 100644 index 00000000000..0823ed13c06 --- /dev/null +++ b/litellm-rust/crates/secrets-types/src/lib.rs @@ -0,0 +1,12 @@ +#![forbid(unsafe_code)] + +mod base_secret_manager; +mod config; +mod error; +mod value; + +pub use base_secret_manager::{BaseSecretManager, async_rotate_secret, validate_secret_name}; +pub use config::{AccessMode, KeyManagementSettings, KeyManagementSystem}; +pub use error::Error; +pub use litellm_auth_types::SecretValue; +pub use value::Secret; diff --git a/litellm-rust/crates/secrets-types/src/value.rs b/litellm-rust/crates/secrets-types/src/value.rs new file mode 100644 index 00000000000..087537fb3eb --- /dev/null +++ b/litellm-rust/crates/secrets-types/src/value.rs @@ -0,0 +1,31 @@ +use crate::SecretValue; + +#[derive(Clone, PartialEq, Eq, veil::Redact)] +pub enum Secret { + String(SecretValue), + Bool(#[redact] bool), + Json(#[redact] serde_json::Value), +} + +impl From for Secret { + fn from(value: SecretValue) -> Self { + Self::String(value) + } +} + +impl Secret { + pub fn from_json(value: serde_json::Value) -> Self { + match value { + serde_json::Value::String(value) => Self::String(SecretValue::new(value)), + serde_json::Value::Bool(value) => Self::Bool(value), + value => Self::Json(value), + } + } + + pub fn as_str(&self) -> Option<&str> { + match self { + Self::String(value) => Some(value.expose()), + Self::Bool(_) | Self::Json(_) => None, + } + } +} diff --git a/litellm-rust/crates/secrets-types/tests/config.rs b/litellm-rust/crates/secrets-types/tests/config.rs new file mode 100644 index 00000000000..4a5f17bc68a --- /dev/null +++ b/litellm-rust/crates/secrets-types/tests/config.rs @@ -0,0 +1,60 @@ +use litellm_secrets_types::{ + AccessMode, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue, +}; +use serde_json::json; + +#[test] +fn config_preserves_defaults_nulls_and_serialized_names() { + let empty: KeyManagementSettings = serde_json::from_value(json!({})).unwrap(); + assert_eq!(empty, KeyManagementSettings::default()); + assert_eq!(empty.access_mode, AccessMode::ReadOnly); + assert_eq!(empty.store_virtual_keys, Some(false)); + assert_eq!(empty.prefix_for_stored_virtual_keys, "litellm/"); + let configured: KeyManagementSettings = serde_json::from_value(json!({ + "hosted_keys": [], "store_virtual_keys": null, "access_mode": "write_only", + "aws_web_identity_token": "private-token", "aws_external_id": "private-id", + "tags": {"stage": "test"}, "replica_regions": ["test-region"] + })) + .unwrap(); + assert!(!configured.access_mode.readable()); + assert_eq!(configured.store_virtual_keys, None); + assert_eq!(configured.hosted_keys.as_deref(), Some([].as_slice())); + assert!(!format!("{configured:?}").contains("private-")); + let serialized = serde_json::to_value(&configured).unwrap(); + assert_eq!(serialized["access_mode"], "write_only"); + assert_eq!(serialized["aws_web_identity_token"], "private-token"); + assert_eq!( + serde_json::from_value::(serialized).unwrap(), + configured + ); +} + +#[rstest::rstest] +#[case::aws_kms("aws_kms", KeyManagementSystem::AwsKms)] +#[case::aws_secret_manager("aws_secret_manager", KeyManagementSystem::AwsSecretManager)] +#[case::google_kms("google_kms", KeyManagementSystem::GoogleKms)] +#[case::google_secret_manager("google_secret_manager", KeyManagementSystem::GoogleSecretManager)] +#[case::azure_key_vault("azure_key_vault", KeyManagementSystem::AzureKeyVault)] +#[case::hashicorp_vault("hashicorp_vault", KeyManagementSystem::HashicorpVault)] +#[case::cyberark("cyberark", KeyManagementSystem::Cyberark)] +#[case::custom("custom", KeyManagementSystem::Custom)] +#[case::local("local", KeyManagementSystem::Local)] +fn key_management_system_serialization_round_trips( + #[case] name: &str, + #[case] system: KeyManagementSystem, +) { + assert_eq!( + serde_json::from_value::(json!(name)).unwrap(), + system + ); + assert_eq!(serde_json::to_value(system).unwrap(), name); +} + +#[test] +fn secret_debug_never_exposes_values() { + assert!( + !format!("{:?}", Secret::String(SecretValue::new("sensitive-value"))) + .contains("sensitive-value") + ); + assert!(!format!("{:?}", Secret::Bool(true)).contains("true")); +} diff --git a/litellm-rust/crates/secrets-types/tests/rotation.rs b/litellm-rust/crates/secrets-types/tests/rotation.rs new file mode 100644 index 00000000000..48a5304ece8 --- /dev/null +++ b/litellm-rust/crates/secrets-types/tests/rotation.rs @@ -0,0 +1,105 @@ +use std::sync::atomic::{AtomicUsize, Ordering}; + +use litellm_secrets_types::{ + BaseSecretManager, Error, SecretValue, async_rotate_secret, validate_secret_name, +}; + +struct Manager { + step: AtomicUsize, + absent_at: Option, +} + +impl BaseSecretManager for Manager { + type Error = Error; + type WriteResponse = &'static str; + type DeleteResponse = (); + + async fn async_read_secret(&self, name: &str) -> Result, Error> { + let step = self.step.fetch_add(1, Ordering::SeqCst); + assert_eq!(name, if step == 0 { "old" } else { "new" }); + Ok((self.absent_at != Some(step)).then(|| SecretValue::new("value"))) + } + + async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + description: Option<&str>, + ) -> Result { + assert_eq!(self.step.fetch_add(1, Ordering::SeqCst), 1); + assert_eq!(name, "new"); + assert_eq!(value.expose(), "replacement"); + assert_eq!(description, Some("Rotated from old")); + Ok("provider-response") + } + + async fn async_delete_secret( + &self, + name: &str, + recovery_window_in_days: i64, + ) -> Result<(), Error> { + assert_eq!(self.step.fetch_add(1, Ordering::SeqCst), 3); + assert_eq!(name, "old"); + assert_eq!(recovery_window_in_days, 7); + Ok(()) + } +} + +#[tokio::test] +async fn rotation_verifies_before_deleting_and_returns_provider_response() { + let manager = Manager { + step: AtomicUsize::new(0), + absent_at: None, + }; + assert_eq!( + async_rotate_secret(&manager, "old", "new", &SecretValue::new("replacement")) + .await + .unwrap(), + "provider-response" + ); + assert_eq!(manager.step.load(Ordering::SeqCst), 4); +} + +#[rstest::rstest] +#[case::current_secret_missing(0, Error::CurrentSecretMissing, 1)] +#[case::new_secret_missing(2, Error::NewSecretMissing, 3)] +#[tokio::test] +async fn missing_old_or_new_value_stops_rotation_before_deletion( + #[case] absent_at: usize, + #[case] expected: Error, + #[case] calls: usize, +) { + let manager = Manager { + step: AtomicUsize::new(0), + absent_at: Some(absent_at), + }; + assert_eq!( + async_rotate_secret(&manager, "old", "new", &SecretValue::new("replacement")) + .await + .unwrap_err(), + expected + ); + assert_eq!(manager.step.load(Ordering::SeqCst), calls); +} + +#[rstest::rstest] +#[case::parent("..")] +#[case::parent_prefix("../x")] +#[case::parent_segment("x/../y")] +#[case::parent_suffix("x/..")] +#[case::line_feed("line\n")] +#[case::next_line("\u{85}")] +#[case::line_separator("\u{2028}")] +#[case::paragraph_separator("\u{2029}")] +fn names_reject_path_traversal_and_control_characters(#[case] name: &str) { + assert_eq!(validate_secret_name(name), Err(Error::UnsafeSecretName)); +} + +#[rstest::rstest] +#[case::embedded_double_dot("release-1.0..2")] +#[case::path_separator("folder/key")] +#[case::empty("")] +#[case::three_dots("...")] +fn names_allow_safe_values(#[case] name: &str) { + assert_eq!(validate_secret_name(name), Ok(())); +} diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml new file mode 100644 index 00000000000..a7e7ec80636 --- /dev/null +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -0,0 +1,34 @@ +[package] +name = "litellm-secrets" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[features] +default = [] +aws = ["dep:litellm-secrets-aws"] +google = ["dep:litellm-secrets-google"] + +[dependencies] +litellm-secrets-types.workspace = true +litellm-secrets-aws = { workspace = true, optional = true } +litellm-secrets-google = { workspace = true, optional = true } +litellm-core-utils.workspace = true +base64.workspace = true +serde.workspace = true +strum.workspace = true +jsonwebtoken.workspace = true +serde_json.workspace = true +thiserror.workspace = true +reqwest.workspace = true +moka.workspace = true +tokio = { workspace = true, features = ["fs"] } + +[dev-dependencies] +rstest.workspace = true +wiremock = "0.6.5" +tempfile = "3" +aws-sdk-kms = "1.120.0" +google-cloud-kms-v1 = "1.14.0" +google-cloud-auth.workspace = true diff --git a/litellm-rust/crates/secrets/README.md b/litellm-rust/crates/secrets/README.md new file mode 100644 index 00000000000..183a39e15bb --- /dev/null +++ b/litellm-rust/crates/secrets/README.md @@ -0,0 +1,11 @@ +# Secret resolution + +Construct `SecretManagerState::new(backend, settings)` for a configured manager or use `SecretManagerState::default()` for environment lookups. The configured backend determines its provider identity. Write-only settings and names excluded by `hosted_keys` use the environment directly. `secret_manager_would_be_consulted` follows the same routing decision as resolution + +`get_secret` returns `Ok(Some(value))` for a found value, `Ok(None)` when no source contains the value, and `Err(error)` when lookup fails. For managed names, resolution checks the manager, then the environment, then the caller's default. An empty string, `false`, or an explicitly stored JSON null is a found value + +Backend failures propagate by default. To allow fallback during a backend failure, construct the resolver with `.with_failure_policy(FailurePolicy::EnvironmentFallback)`. It then tries the environment and default, in that order. If neither exists, the original error is returned. This policy applies to manager lookups. Explicit OIDC references retain their own authentication errors and never fall back to environment secrets under the reference name + +`get_secret` preserves value types. `get_secret_str` accepts a string default and rejects boolean or JSON values with `Error::TypeMismatch`. `get_secret_bool` accepts a boolean default and converts strings containing `true` or `false`, ignoring surrounding whitespace and ASCII case. Other strings and JSON values produce `Error::TypeMismatch`. Conversion failures never activate fallback or replace a found value with the default + +Provider payloads remain strings unless explicitly selecting a field from an AWS primary JSON secret. Google caches only successfully decoded string payloads, so reads have identical values and types before and after caching. Confirmed absence and failed reads are not cached. AWS resource-not-found responses and Google HTTP 404 responses indicate absence. Other provider errors remain errors, and successful responses without the required payload are malformed responses rather than missing secrets diff --git a/litellm-rust/crates/secrets/src/error.rs b/litellm-rust/crates/secrets/src/error.rs new file mode 100644 index 00000000000..0c6e681b8aa --- /dev/null +++ b/litellm-rust/crates/secrets/src/error.rs @@ -0,0 +1,33 @@ +#[derive(Debug, thiserror::Error)] +pub enum Error { + #[error("encrypted environment value is missing")] + MissingCiphertext, + #[error("ciphertext is not valid base64 for the configured manager")] + InvalidCiphertext, + #[error("decrypted value is not UTF-8")] + Utf8, + #[error("unsupported OIDC provider or missing build feature")] + UnsupportedOidc, + #[error("OIDC reference requires a provider and audience")] + InvalidOidc, + #[error("OIDC environment variable is missing")] + MissingEnvironment, + #[error("OIDC request failed")] + OidcHttp, + #[error("OIDC provider returned HTTP {0}")] + OidcStatus(u16), + #[error("OIDC response is invalid")] + OidcResponse, + #[error("OIDC file path must be absolute and within the credential allowlist")] + UnsafeOidcPath, + #[error("OIDC file could not be read")] + OidcFile, + #[error("secret cannot be converted to {expected}")] + TypeMismatch { expected: &'static str }, + #[cfg(feature = "aws")] + #[error(transparent)] + Aws(#[from] litellm_secrets_aws::Error), + #[cfg(feature = "google")] + #[error(transparent)] + Google(#[from] litellm_secrets_google::Error), +} diff --git a/litellm-rust/crates/secrets/src/handler.rs b/litellm-rust/crates/secrets/src/handler.rs new file mode 100644 index 00000000000..943ffdf6158 --- /dev/null +++ b/litellm-rust/crates/secrets/src/handler.rs @@ -0,0 +1,117 @@ +use litellm_core_utils::settings::Lookup; + +use crate::{Error, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue}; + +#[derive(Clone)] +pub enum SecretManager { + Local, + #[cfg(feature = "aws")] + AwsKms(crate::aws::AwsKms), + #[cfg(feature = "aws")] + AwsSecretsManagerV2(crate::aws::AwsSecretsManagerV2), + #[cfg(feature = "google")] + GoogleKms(crate::google::GoogleKms), + #[cfg(feature = "google")] + GoogleSecretManager(crate::google::GoogleSecretManager), +} + +impl SecretManager { + pub fn system(&self) -> KeyManagementSystem { + match self { + Self::Local => KeyManagementSystem::Local, + #[cfg(feature = "aws")] + Self::AwsKms(_) => KeyManagementSystem::AwsKms, + #[cfg(feature = "aws")] + Self::AwsSecretsManagerV2(_) => KeyManagementSystem::AwsSecretManager, + #[cfg(feature = "google")] + Self::GoogleKms(_) => KeyManagementSystem::GoogleKms, + #[cfg(feature = "google")] + Self::GoogleSecretManager(_) => KeyManagementSystem::GoogleSecretManager, + } + } +} + +pub async fn get_secret_from_manager( + client: &SecretManager, + secret_name: &str, + _settings: &KeyManagementSettings, + environment: &(dyn Lookup + Send + Sync), +) -> Result, Error> { + match client { + SecretManager::Local => Ok(environment + .get(secret_name) + .map(SecretValue::new) + .map(Secret::String)), + #[cfg(feature = "aws")] + SecretManager::AwsKms(client) => { + let ciphertext = environment + .get(secret_name) + .ok_or(Error::MissingCiphertext)?; + let plaintext = client + .decrypt(decode_ciphertext(&ciphertext, Base64Mode::Permissive)?) + .await?; + let value = String::from_utf8(plaintext).map_err(|_| Error::Utf8)?; + Ok(Some(Secret::String(SecretValue::new(value.trim())))) + } + #[cfg(feature = "google")] + SecretManager::GoogleKms(client) => { + let ciphertext = environment + .get(secret_name) + .ok_or(Error::MissingCiphertext)?; + let plaintext = client + .decrypt(decode_ciphertext(&ciphertext, Base64Mode::Canonical)?) + .await?; + let value = String::from_utf8(plaintext).map_err(|_| Error::Utf8)?; + Ok(Some(Secret::String(SecretValue::new(value)))) + } + #[cfg(feature = "aws")] + SecretManager::AwsSecretsManagerV2(client) => client + .read_secret_for_resolver( + secret_name, + _settings.primary_secret_name.as_deref(), + environment, + ) + .await + .map_err(Error::from), + #[cfg(feature = "google")] + SecretManager::GoogleSecretManager(client) => client + .get_secret_from_google_secret_manager(secret_name) + .await + .map_err(Error::from), + } +} + +#[cfg(any(feature = "aws", feature = "google"))] +#[derive(Clone, Copy)] +enum Base64Mode { + #[cfg(feature = "google")] + Canonical, + #[cfg(feature = "aws")] + Permissive, +} + +#[cfg(any(feature = "aws", feature = "google"))] +fn decode_ciphertext(value: &str, mode: Base64Mode) -> Result, Error> { + use base64::{Engine, engine::general_purpose::STANDARD}; + let canonical = match mode { + #[cfg(feature = "google")] + Base64Mode::Canonical => true, + #[cfg(feature = "aws")] + Base64Mode::Permissive => false, + }; + let encoded = if canonical { + value.to_owned() + } else { + value + .chars() + .filter(|c| c.is_ascii_alphanumeric() || matches!(c, '+' | '/' | '=')) + .collect() + }; + let ciphertext = STANDARD + .decode(&encoded) + .map_err(|_| Error::InvalidCiphertext)?; + if canonical && STANDARD.encode(&ciphertext) != encoded { + return Err(Error::InvalidCiphertext); + } + Ok(ciphertext) +} diff --git a/litellm-rust/crates/secrets/src/lib.rs b/litellm-rust/crates/secrets/src/lib.rs new file mode 100644 index 00000000000..ff2e95f7b2f --- /dev/null +++ b/litellm-rust/crates/secrets/src/lib.rs @@ -0,0 +1,21 @@ +#![forbid(unsafe_code)] + +mod error; +mod handler; +mod oidc; +mod resolver; +mod state; + +pub use error::Error; +pub use handler::{SecretManager, get_secret_from_manager}; +pub use litellm_secrets_types::{ + AccessMode, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue, +}; +pub use oidc::{OidcProvider, OidcReference, OidcResolver}; +pub use resolver::{FailurePolicy, SecretResolver}; +pub use state::{SecretManagerState, secret_manager_would_be_consulted}; + +#[cfg(feature = "aws")] +pub use litellm_secrets_aws as aws; +#[cfg(feature = "google")] +pub use litellm_secrets_google as google; diff --git a/litellm-rust/crates/secrets/src/oidc.rs b/litellm-rust/crates/secrets/src/oidc.rs new file mode 100644 index 00000000000..fd477859bf6 --- /dev/null +++ b/litellm-rust/crates/secrets/src/oidc.rs @@ -0,0 +1,269 @@ +use std::{ + path::Path, + time::{Duration, SystemTime, UNIX_EPOCH}, +}; + +use jsonwebtoken::dangerous::insecure_decode_claims; +use litellm_core_utils::settings::Lookup; +use moka::future::Cache; +use serde::Deserialize; + +use crate::{Error, SecretValue}; + +const GOOGLE_TOKEN_MAX_TTL: Duration = Duration::from_secs(3540); +const GITHUB_TOKEN_TTL: Duration = Duration::from_secs(295); +const TOKEN_EXPIRY_MARGIN_SECONDS: f64 = 60.0; +const CIRCLE_OIDC_TOKEN: &str = "CIRCLE_OIDC_TOKEN"; +const CIRCLE_OIDC_TOKEN_V2: &str = "CIRCLE_OIDC_TOKEN_V2"; +const AZURE_FEDERATED_TOKEN_FILE: &str = "AZURE_FEDERATED_TOKEN_FILE"; +const ACTIONS_ID_TOKEN_REQUEST_URL: &str = "ACTIONS_ID_TOKEN_REQUEST_URL"; +const ACTIONS_ID_TOKEN_REQUEST_TOKEN: &str = "ACTIONS_ID_TOKEN_REQUEST_TOKEN"; +const OIDC_ALLOWED_CREDENTIAL_DIRS: &str = "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS"; +const DEFAULT_CREDENTIAL_DIRS: &str = "/var/run/secrets,/run/secrets"; + +#[derive(Clone, Copy, Debug, Eq, Hash, PartialEq, strum::EnumString, strum::AsRefStr)] +#[strum(serialize_all = "snake_case")] +pub enum OidcProvider { + Google, + #[strum(serialize = "circleci")] + CircleCi, + #[strum(serialize = "circleci_v2")] + CircleCiV2, + Github, + Azure, + File, + Env, + EnvPath, +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub struct OidcReference<'a> { + pub provider: OidcProvider, + pub audience: &'a str, +} + +impl<'a> TryFrom<&'a str> for OidcReference<'a> { + type Error = Error; + + fn try_from(reference: &'a str) -> Result { + let (provider, audience) = reference + .strip_prefix("oidc/") + .and_then(|body| body.split_once('/')) + .ok_or(Error::InvalidOidc)?; + Ok(Self { + provider: provider.parse().map_err(|_| Error::UnsupportedOidc)?, + audience, + }) + } +} + +#[derive(Deserialize)] +struct OidcTokenClaims { + exp: Option, +} + +#[derive(Deserialize)] +#[serde(untagged)] +enum NumericDate { + Number(f64), + String(String), +} + +impl NumericDate { + fn seconds(self) -> Option { + match self { + Self::Number(value) => Some(value), + Self::String(value) => value.parse().ok(), + } + .filter(|value| value.is_finite()) + } +} + +pub struct OidcResolver { + client: reqwest::Client, + google_identity_endpoint: reqwest::Url, + cache: Cache, + clock: fn() -> SystemTime, +} + +impl Default for OidcResolver { + fn default() -> Self { + let client = reqwest::Client::builder() + .timeout(Duration::from_secs(600)) + .connect_timeout(Duration::from_secs(5)) + .build() + .expect("HTTP client configuration"); + Self::new( + client, + reqwest::Url::parse("http://metadata.google.internal/computeMetadata/v1/instance/service-accounts/default/identity").expect("static URL"), + ) + } +} + +impl OidcResolver { + pub fn new(client: reqwest::Client, google_identity_endpoint: reqwest::Url) -> Self { + Self { + client, + google_identity_endpoint, + cache: Cache::builder() + .max_capacity(200) + .time_to_live(GOOGLE_TOKEN_MAX_TTL) + .build(), + clock: SystemTime::now, + } + } + + pub fn with_clock(self, clock: fn() -> SystemTime) -> Self { + Self { clock, ..self } + } + + pub async fn resolve( + &self, + reference: &str, + environment: &(dyn Lookup + Send + Sync), + ) -> Result, Error> { + let OidcReference { provider, audience } = reference.try_into()?; + match provider { + OidcProvider::CircleCi => required_env(environment, CIRCLE_OIDC_TOKEN) + .map(SecretValue::new) + .map(Some), + OidcProvider::CircleCiV2 => required_env(environment, CIRCLE_OIDC_TOKEN_V2) + .map(SecretValue::new) + .map(Some), + OidcProvider::Env => required_env(environment, audience) + .map(SecretValue::new) + .map(Some), + OidcProvider::EnvPath => read_file(&required_env(environment, audience)?) + .await + .map(Some), + OidcProvider::File => read_allowed_file(audience, environment).await.map(Some), + OidcProvider::Azure => { + if let Some(path) = environment.get(AZURE_FEDERATED_TOKEN_FILE) { + return read_file(&path).await.map(Some); + } + Err(Error::UnsupportedOidc) + } + OidcProvider::Github => { + let url = required_env(environment, ACTIONS_ID_TOKEN_REQUEST_URL)?; + let authorization = required_env(environment, ACTIONS_ID_TOKEN_REQUEST_TOKEN)?; + if let Some(value) = self.cached(reference).await { + return Ok(Some(value)); + } + let response = self + .client + .get(url) + .query(&[("audience", audience)]) + .bearer_auth(authorization) + .header("Accept", "application/json; api-version=2.0") + .send() + .await + .map_err(|_| Error::OidcHttp)?; + if response.status() != reqwest::StatusCode::OK { + return Err(Error::OidcStatus(response.status().as_u16())); + } + #[derive(Deserialize)] + struct Token { + value: Option, + } + let token: Token = response.json().await.map_err(|_| Error::OidcResponse)?; + if let Some(value) = &token.value { + self.cache + .insert( + reference.to_owned(), + (value.clone(), (self.clock)() + GITHUB_TOKEN_TTL), + ) + .await; + } + Ok(token.value) + } + OidcProvider::Google => { + if !cfg!(feature = "google") { + return Err(Error::UnsupportedOidc); + } + if let Some(value) = self.cached(reference).await { + return Ok(Some(value)); + } + let response = self + .client + .get(self.google_identity_endpoint.clone()) + .query(&[("audience", audience)]) + .header("Metadata-Flavor", "Google") + .send() + .await + .map_err(|_| Error::OidcHttp)?; + if response.status() != reqwest::StatusCode::OK { + return Err(Error::OidcStatus(response.status().as_u16())); + } + let token = response.text().await.map_err(|_| Error::OidcResponse)?; + let now = (self.clock)(); + let ttl = oidc_token_cache_ttl(&token, now, GOOGLE_TOKEN_MAX_TTL); + let value = SecretValue::new(token); + if let Some(ttl) = ttl.filter(|ttl| !ttl.is_zero()) { + self.cache + .insert(reference.to_owned(), (value.clone(), now + ttl)) + .await; + } + Ok(Some(value)) + } + } + } + + async fn cached(&self, reference: &str) -> Option { + self.cache + .get(reference) + .await + .and_then(|(value, expires)| ((self.clock)() < expires).then_some(value)) + } +} + +fn required_env(environment: &dyn Lookup, name: &str) -> Result { + environment.get(name).ok_or(Error::MissingEnvironment) +} + +async fn read_file(path: &str) -> Result { + tokio::fs::read_to_string(path) + .await + .map(|value| SecretValue::new(value.replace("\r\n", "\n").replace('\r', "\n"))) + .map_err(|_| Error::OidcFile) +} + +async fn read_allowed_file( + path: &str, + environment: &(dyn Lookup + Sync), +) -> Result { + if !Path::new(path).is_absolute() { + return Err(Error::UnsafeOidcPath); + } + let resolved = tokio::fs::canonicalize(path) + .await + .map_err(|_| Error::OidcFile)?; + let allowed = environment + .get(OIDC_ALLOWED_CREDENTIAL_DIRS) + .filter(|v| !v.is_empty()) + .unwrap_or_else(|| DEFAULT_CREDENTIAL_DIRS.into()); + for directory in allowed.split(',').map(str::trim).filter(|d| !d.is_empty()) { + if let Ok(directory) = tokio::fs::canonicalize(directory).await + && resolved.starts_with(directory) + { + return tokio::fs::read_to_string(&resolved) + .await + .map(|value| SecretValue::new(value.replace("\r\n", "\n").replace('\r', "\n"))) + .map_err(|_| Error::OidcFile); + } + } + Err(Error::UnsafeOidcPath) +} + +fn oidc_token_cache_ttl(token: &str, now: SystemTime, max_ttl: Duration) -> Option { + let fallback = Some(max_ttl); + let Ok(claims) = insecure_decode_claims::(token) else { + return fallback; + }; + let Some(exp) = claims.exp.and_then(NumericDate::seconds) else { + return fallback; + }; + let seconds = exp.trunc() + - now.duration_since(UNIX_EPOCH).ok()?.as_secs() as f64 + - TOKEN_EXPIRY_MARGIN_SECONDS; + (seconds > 0.0).then(|| Duration::from_secs_f64(seconds.min(max_ttl.as_secs_f64()))) +} diff --git a/litellm-rust/crates/secrets/src/resolver.rs b/litellm-rust/crates/secrets/src/resolver.rs new file mode 100644 index 00000000000..89439893852 --- /dev/null +++ b/litellm-rust/crates/secrets/src/resolver.rs @@ -0,0 +1,135 @@ +use std::sync::Arc; + +use litellm_core_utils::settings::{Lookup, ProcessEnvironment}; + +use crate::state::{LookupTarget, normalize_secret_name}; +use crate::{Error, OidcResolver, Secret, SecretManagerState, SecretValue}; + +#[derive(Clone, Copy, Debug, Default, PartialEq, Eq)] +pub enum FailurePolicy { + #[default] + Propagate, + EnvironmentFallback, +} + +pub struct SecretResolver { + state: Arc, + environment: Arc, + oidc: OidcResolver, + failure_policy: FailurePolicy, +} + +impl Default for SecretResolver { + fn default() -> Self { + Self::new( + Arc::new(SecretManagerState::default()), + Arc::new(ProcessEnvironment), + OidcResolver::default(), + ) + } +} + +impl SecretResolver { + pub fn new( + state: Arc, + environment: Arc, + oidc: OidcResolver, + ) -> Self { + Self { + state, + environment, + oidc, + failure_policy: FailurePolicy::default(), + } + } + + pub fn with_failure_policy(self, failure_policy: FailurePolicy) -> Self { + Self { + failure_policy, + ..self + } + } + + pub async fn get_secret( + &self, + name: &str, + default_value: Option, + ) -> Result, Error> { + let name = normalize_secret_name(name); + if name.starts_with("oidc/") { + return self + .oidc + .resolve(name, self.environment.as_ref()) + .await + .map(|value| value.map(Secret::String).or(default_value)); + } + let LookupTarget::Manager { backend, settings } = self.state.lookup_target(name) else { + return Ok(self.environment_secret(name).or(default_value)); + }; + match crate::get_secret_from_manager(backend, name, settings, self.environment.as_ref()) + .await + { + Ok(value) => Ok(value + .or_else(|| self.environment_secret(name)) + .or(default_value)), + Err(error) => match self.failure_policy { + FailurePolicy::Propagate => Err(error), + FailurePolicy::EnvironmentFallback => self + .environment_secret(name) + .or(default_value) + .map(Some) + .ok_or(error), + }, + } + } + + fn environment_secret(&self, name: &str) -> Option { + self.environment + .get(name) + .map(SecretValue::new) + .map(Secret::String) + } + + pub async fn get_secret_str( + &self, + name: &str, + default_value: Option, + ) -> Result, Error> { + match self + .get_secret(name, default_value.map(Secret::String)) + .await? + { + Some(Secret::String(value)) => Ok(Some(value)), + None => Ok(None), + Some(Secret::Bool(_) | Secret::Json(_)) => { + Err(Error::TypeMismatch { expected: "string" }) + } + } + } + + pub async fn get_secret_bool( + &self, + name: &str, + default_value: Option, + ) -> Result, Error> { + match self + .get_secret(name, default_value.map(Secret::Bool)) + .await? + { + Some(Secret::Bool(value)) => Ok(Some(value)), + Some(Secret::String(value)) => { + match value.expose().trim().to_ascii_lowercase().as_str() { + "true" => Ok(Some(true)), + "false" => Ok(Some(false)), + _ => Err(Error::TypeMismatch { + expected: "boolean", + }), + } + } + Some(Secret::Json(_)) => Err(Error::TypeMismatch { + expected: "boolean", + }), + None => Ok(None), + } + } +} diff --git a/litellm-rust/crates/secrets/src/state.rs b/litellm-rust/crates/secrets/src/state.rs new file mode 100644 index 00000000000..7854d763ef2 --- /dev/null +++ b/litellm-rust/crates/secrets/src/state.rs @@ -0,0 +1,59 @@ +use crate::{KeyManagementSettings, KeyManagementSystem, SecretManager}; + +pub(crate) enum LookupTarget<'a> { + Environment, + Manager { + backend: &'a SecretManager, + settings: &'a KeyManagementSettings, + }, +} + +pub(crate) fn normalize_secret_name(name: &str) -> &str { + name.strip_prefix("os.environ/").unwrap_or(name) +} + +#[derive(Clone, Default)] +pub struct SecretManagerState { + manager: Option<(SecretManager, KeyManagementSettings)>, +} + +impl SecretManagerState { + pub fn new(backend: SecretManager, settings: KeyManagementSettings) -> Self { + Self { + manager: Some((backend, settings)), + } + } + + pub fn system(&self) -> Option { + self.backend().map(SecretManager::system) + } + + pub fn settings(&self) -> Option<&KeyManagementSettings> { + self.manager.as_ref().map(|(_, settings)| settings) + } + + pub fn backend(&self) -> Option<&SecretManager> { + self.manager.as_ref().map(|(backend, _)| backend) + } + + pub(crate) fn lookup_target(&self, name: &str) -> LookupTarget<'_> { + match &self.manager { + Some((backend, settings)) + if backend.system() != KeyManagementSystem::Local + && settings.access_mode.readable() + && settings + .hosted_keys + .as_ref() + .is_none_or(|keys| keys.iter().any(|key| key == name)) => + { + LookupTarget::Manager { backend, settings } + } + _ => LookupTarget::Environment, + } + } +} + +pub fn secret_manager_would_be_consulted(state: &SecretManagerState, name: &str) -> bool { + let name = normalize_secret_name(name); + !name.starts_with("oidc/") && matches!(state.lookup_target(name), LookupTarget::Manager { .. }) +} diff --git a/litellm-rust/crates/secrets/tests/handler.rs b/litellm-rust/crates/secrets/tests/handler.rs new file mode 100644 index 00000000000..a2cbbd843e1 --- /dev/null +++ b/litellm-rust/crates/secrets/tests/handler.rs @@ -0,0 +1,107 @@ +#[cfg(feature = "aws")] +#[tokio::test] +async fn aws_handler_reads_ciphertext_decodes_trims_and_redacts() { + use aws_sdk_kms::{ + Client, + config::{BehaviorVersion, Credentials, Region}, + }; + use base64::{Engine, engine::general_purpose::STANDARD}; + use litellm_secrets::{ + Error, KeyManagementSettings, SecretManager, aws::AwsKms, get_secret_from_manager, + }; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::body_json}; + + let server = MockServer::start().await; + Mock::given(body_json( + serde_json::json!({"CiphertextBlob": STANDARD.encode("encrypted")}), + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"Plaintext":STANDARD.encode(" value\n")})), + ) + .expect(1) + .mount(&server) + .await; + let client = Client::from_conf( + aws_sdk_kms::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new("us-east-1")) + .credentials_provider(Credentials::new("test", "test", None, None, "test")) + .endpoint_url(server.uri()) + .build(), + ); + let manager = SecretManager::AwsKms(AwsKms::new(client)); + let settings = KeyManagementSettings::default(); + let value = get_secret_from_manager(&manager, "KEY", &settings, &|name: &str| { + assert_eq!(name, "KEY"); + Some(format!(" {}\n", STANDARD.encode("encrypted"))) + }) + .await + .unwrap() + .unwrap(); + assert_eq!(value.as_str(), Some("value")); + assert!(!format!("{value:?}").contains("value")); + assert!(matches!( + get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None).await, + Err(Error::MissingCiphertext) + )); + assert!(matches!( + get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| Some("abc".into())).await, + Err(Error::InvalidCiphertext) + )); +} + +#[cfg(feature = "google")] +#[tokio::test] +async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whitespace() { + use base64::{Engine, engine::general_purpose::STANDARD}; + use google_cloud_kms_v1::client::KeyManagementService; + use litellm_secrets::{ + Error, KeyManagementSettings, SecretManager, get_secret_from_manager, google::GoogleKms, + }; + use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_json, path}, + }; + + let server = MockServer::start().await; + let resource = "projects/project/locations/global/keyRings/ring/cryptoKeys/key"; + Mock::given(path(format!("/v1/{resource}:decrypt"))) + .and(body_json( + serde_json::json!({"ciphertext":STANDARD.encode("encrypted")}), + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"plaintext":STANDARD.encode(" value\n")})), + ) + .expect(1) + .mount(&server) + .await; + let client = KeyManagementService::builder() + .with_endpoint(server.uri()) + .with_credentials(google_cloud_auth::credentials::anonymous::Builder::new().build()) + .build() + .await + .unwrap(); + let manager = SecretManager::GoogleKms(GoogleKms::new(client, resource.into())); + let settings = KeyManagementSettings::default(); + let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| { + Some(STANDARD.encode("encrypted")) + }) + .await + .unwrap() + .unwrap(); + assert_eq!(value.as_str(), Some(" value\n")); + assert!(matches!( + get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| Some(format!( + " {}", + STANDARD.encode("encrypted") + ))) + .await, + Err(Error::InvalidCiphertext) + )); + assert!(matches!( + get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None).await, + Err(Error::MissingCiphertext) + )); +} diff --git a/litellm-rust/crates/secrets/tests/oidc.rs b/litellm-rust/crates/secrets/tests/oidc.rs new file mode 100644 index 00000000000..b17e7de7f9d --- /dev/null +++ b/litellm-rust/crates/secrets/tests/oidc.rs @@ -0,0 +1,295 @@ +use std::{collections::BTreeMap, sync::Arc}; + +use litellm_core_utils::settings::Lookup; +use litellm_secrets::{Error, OidcResolver, Secret, SecretManagerState, SecretResolver}; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{header, method, path, query_param}, +}; + +fn environment(pairs: &[(&str, &str)]) -> Arc { + let values: BTreeMap = pairs + .iter() + .map(|(k, v)| (k.to_string(), v.to_string())) + .collect(); + Arc::new(move |name: &str| values.get(name).cloned()) +} + +#[rstest::rstest] +#[case::environment("oidc/env/TOKEN", "true")] +#[case::circleci("oidc/circleci/audience", "circle")] +#[case::circleci_v2("oidc/circleci_v2/audience", "circle-v2")] +#[tokio::test] +async fn environment_sources_resolve_expected_value( + #[case] reference: &str, + #[case] expected: &str, +) { + let env = environment(&[ + ("TOKEN", "true"), + ("CIRCLE_OIDC_TOKEN", "circle"), + ("CIRCLE_OIDC_TOKEN_V2", "circle-v2"), + ]); + assert_eq!( + OidcResolver::default() + .resolve(reference, env.as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + expected + ); +} + +#[tokio::test] +async fn environment_sources_bypass_boolean_conversion_and_defaults() { + let env = environment(&[("TOKEN", "true")]); + let oidc = OidcResolver::default(); + let resolver = SecretResolver::new(Arc::new(SecretManagerState::default()), env, oidc); + assert_eq!( + resolver + .get_secret_str("os.environ/oidc/env/TOKEN", None) + .await + .unwrap() + .unwrap() + .expose(), + "true" + ); + assert_eq!( + resolver + .get_secret_bool("oidc/env/TOKEN", None) + .await + .unwrap(), + Some(true) + ); + assert!(matches!( + resolver + .get_secret("oidc/env/MISSING", Some(Secret::Bool(true))) + .await, + Err(Error::MissingEnvironment) + )); + assert!(matches!( + resolver.get_secret("oidc/invalid", None).await, + Err(Error::InvalidOidc) + )); +} + +#[tokio::test] +async fn github_requests_are_authenticated_cached_and_revalidate_environment() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/token")) + .and(query_param("audience", "https://service/oidc/path")) + .and(header("authorization", "Bearer request-token")) + .and(header("accept", "application/json; api-version=2.0")) + .respond_with( + ResponseTemplate::new(200).set_body_json(serde_json::json!({"value":"identity-token"})), + ) + .expect(1) + .mount(&server) + .await; + let env = environment(&[ + ( + "ACTIONS_ID_TOKEN_REQUEST_URL", + &format!("{}/token", server.uri()), + ), + ("ACTIONS_ID_TOKEN_REQUEST_TOKEN", "request-token"), + ]); + let oidc = OidcResolver::default(); + for _ in 0..2 { + assert_eq!( + oidc.resolve("oidc/github/https://service/oidc/path", env.as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + "identity-token" + ); + } + assert!(matches!( + oidc.resolve( + "oidc/github/https://service/oidc/path", + environment(&[]).as_ref() + ) + .await, + Err(Error::MissingEnvironment) + )); +} + +#[tokio::test] +async fn file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit() { + let allowed = tempfile::tempdir().unwrap(); + let outside = tempfile::tempdir().unwrap(); + let token = allowed.path().join("token"); + let private = outside.path().join("private"); + std::fs::write(&token, "token\r\n").unwrap(); + std::fs::write(&private, "outside").unwrap(); + let env = environment(&[ + ( + "LITELLM_OIDC_ALLOWED_CREDENTIAL_DIRS", + allowed.path().to_str().unwrap(), + ), + ("PATH_TOKEN", private.to_str().unwrap()), + ("AZURE_FEDERATED_TOKEN_FILE", token.to_str().unwrap()), + ]); + let oidc = OidcResolver::default(); + assert_eq!( + oidc.resolve(&format!("oidc/file/{}", token.display()), env.as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + "token\n" + ); + assert!(matches!( + oidc.resolve("oidc/file/relative", env.as_ref()).await, + Err(Error::UnsafeOidcPath) + )); + assert!(matches!( + oidc.resolve(&format!("oidc/file/{}", private.display()), env.as_ref()) + .await, + Err(Error::UnsafeOidcPath) + )); + assert_eq!( + oidc.resolve("oidc/env_path/PATH_TOKEN", env.as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + "outside" + ); + assert_eq!( + oidc.resolve("oidc/azure/scope", env.as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + "token\n" + ); + #[cfg(unix)] + { + let link = allowed.path().join("link"); + std::os::unix::fs::symlink(&private, &link).unwrap(); + assert!(matches!( + oidc.resolve(&format!("oidc/file/{}", link.display()), env.as_ref()) + .await, + Err(Error::UnsafeOidcPath) + )); + } +} + +#[cfg(feature = "google")] +#[rstest::rstest] +#[case::at_refresh_boundary(serde_json::json!(1060), 2)] +#[case::beyond_refresh_boundary(serde_json::json!(1061), 1)] +#[case::already_expired(serde_json::json!(999), 2)] +#[case::string_expiry(serde_json::json!("999"), 2)] +#[case::fractional_expiry(serde_json::json!(1060.9), 2)] +#[case::negative_expiry(serde_json::json!(-1), 2)] +#[case::null_expiry(serde_json::Value::Null, 1)] +#[case::unreadable_expiry(serde_json::json!("invalid"), 1)] +#[case::nonfinite_expiry(serde_json::json!("NaN"), 1)] +#[tokio::test] +async fn google_expiry_caps_cache_and_preserves_audience( + #[case] expiry: serde_json::Value, + #[case] calls: u64, +) { + use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; + use std::time::{Duration, SystemTime, UNIX_EPOCH}; + fn now() -> SystemTime { + UNIX_EPOCH + Duration::from_secs(1000) + } + let server = MockServer::start().await; + let token = format!( + "{}.{}.signature", + URL_SAFE_NO_PAD.encode(serde_json::json!({"alg":"RS256","typ":"JWT"}).to_string()), + URL_SAFE_NO_PAD.encode(serde_json::json!({"exp":expiry}).to_string()) + ); + Mock::given(method("GET")) + .and(header("metadata-flavor", "Google")) + .and(query_param("audience", "https://service/oidc/path")) + .respond_with(ResponseTemplate::new(200).set_body_string(&token)) + .expect(calls) + .mount(&server) + .await; + let oidc = + OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()).with_clock(now); + for _ in 0..2 { + assert_eq!( + oidc.resolve( + "oidc/google/https://service/oidc/path", + environment(&[]).as_ref() + ) + .await + .unwrap() + .unwrap() + .expose(), + token + ); + } +} + +#[cfg(not(feature = "google"))] +#[tokio::test] +async fn google_oidc_requires_its_build_feature() { + assert!(matches!( + OidcResolver::default() + .resolve("oidc/google/audience", environment(&[]).as_ref()) + .await, + Err(Error::UnsupportedOidc) + )); +} + +#[tokio::test] +async fn azure_oidc_without_a_token_file_requires_an_unimplemented_backend() { + assert!(matches!( + OidcResolver::default() + .resolve("oidc/azure/scope", environment(&[]).as_ref()) + .await, + Err(Error::UnsupportedOidc) + )); +} + +#[rstest::rstest] +#[case::missing_prefix("env/TOKEN", false)] +#[case::missing_audience_separator("oidc/env", false)] +#[case::unknown_provider("oidc/unknown/TOKEN", true)] +#[tokio::test] +async fn invalid_references_fail_before_environment_lookup( + #[case] reference: &str, + #[case] unsupported: bool, +) { + let error = OidcResolver::default() + .resolve(reference, &|_: &str| { + panic!("invalid reference reached environment lookup") + }) + .await + .unwrap_err(); + assert!(matches!(error, Error::UnsupportedOidc) == unsupported); + assert!(matches!(error, Error::InvalidOidc) != unsupported); +} + +#[cfg(feature = "google")] +#[rstest::rstest] +#[case::opaque("opaque-token")] +#[case::missing_expiry("header.e30.signature")] +#[tokio::test] +async fn unreadable_expiry_keeps_python_cache_fallback(#[case] token: &str) { + let server = MockServer::start().await; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(200).set_body_string(token)) + .expect(1) + .mount(&server) + .await; + let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()); + for _ in 0..2 { + assert_eq!( + resolver + .resolve("oidc/google/audience", environment(&[]).as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + token, + ); + } +} diff --git a/litellm-rust/crates/secrets/tests/resolution.rs b/litellm-rust/crates/secrets/tests/resolution.rs new file mode 100644 index 00000000000..3a826092d72 --- /dev/null +++ b/litellm-rust/crates/secrets/tests/resolution.rs @@ -0,0 +1,368 @@ +use std::sync::Arc; + +use litellm_secrets::{ + Error, KeyManagementSettings, OidcResolver, Secret, SecretManager, SecretManagerState, + SecretResolver, SecretValue, secret_manager_would_be_consulted, +}; + +fn resolver(value: Option<&str>, configured: bool) -> SecretResolver { + let state = if configured { + SecretManagerState::new(SecretManager::Local, KeyManagementSettings::default()) + } else { + SecretManagerState::default() + }; + let value = value.map(str::to_owned); + SecretResolver::new( + Arc::new(state), + Arc::new(move |_: &str| value.clone()), + OidcResolver::default(), + ) +} + +#[rstest::rstest] +#[case("true", Some(true))] +#[case(" FALSE ", Some(false))] +#[case("(True)", None)] +#[case("False # comment", None)] +#[case("1", None)] +#[case("secret", None)] +#[tokio::test] +async fn conversion_is_explicit_and_independent_of_manager_configuration( + #[case] input: &str, + #[case] boolean: Option, + #[values(false, true)] configured: bool, +) { + let resolver = resolver(Some(input), configured); + assert_eq!( + resolver.get_secret("key", None).await.unwrap(), + Some(Secret::String(SecretValue::new(input))) + ); + assert_eq!( + resolver + .get_secret_str("key", None) + .await + .unwrap() + .unwrap() + .expose(), + input + ); + match boolean { + Some(value) => assert_eq!( + resolver.get_secret_bool("key", None).await.unwrap(), + Some(value) + ), + None => assert!(matches!( + resolver.get_secret_bool("key", Some(true)).await, + Err(Error::TypeMismatch { + expected: "boolean" + }) + )), + } +} + +#[rstest::rstest] +#[tokio::test] +async fn defaults_apply_only_to_absence(#[values(false, true)] configured: bool) { + let missing = resolver(None, configured); + assert_eq!(missing.get_secret("key", None).await.unwrap(), None); + assert_eq!( + missing.get_secret_bool("key", Some(false)).await.unwrap(), + Some(false) + ); + assert_eq!( + missing + .get_secret_str("key", Some(SecretValue::new("default"))) + .await + .unwrap() + .unwrap() + .expose(), + "default" + ); + for value in [ + Secret::Bool(false), + Secret::from_json(serde_json::json!({"key":1})), + Secret::from_json(serde_json::Value::Null), + ] { + assert_eq!( + missing + .get_secret("key", Some(value.clone())) + .await + .unwrap(), + Some(value) + ); + } + assert_eq!( + resolver(Some(""), configured) + .get_secret_str("key", Some(SecretValue::new("default"))) + .await + .unwrap() + .unwrap() + .expose(), + "" + ); +} + +#[tokio::test] +async fn prefix_is_removed_once_and_local_manager_is_not_consulted() { + let state = SecretManagerState::new(SecretManager::Local, KeyManagementSettings::default()); + assert_eq!( + state.system(), + Some(litellm_secrets::KeyManagementSystem::Local) + ); + assert!(!secret_manager_would_be_consulted( + &state, + "os.environ/os.environ/KEY" + )); + let resolver = SecretResolver::new( + Arc::new(state), + Arc::new(|name: &str| (name == "os.environ/KEY").then(|| "value".into())), + OidcResolver::default(), + ); + assert_eq!( + resolver + .get_secret_str("os.environ/os.environ/KEY", None) + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[tokio::test] +async fn resolver_future_can_run_on_a_tokio_worker() { + let resolver = resolver(Some("worker-value"), false); + let result = tokio::spawn(async move { resolver.get_secret_str("KEY", None).await }) + .await + .unwrap() + .unwrap(); + assert_eq!(result.unwrap().expose(), "worker-value"); +} + +#[cfg(feature = "aws")] +mod aws { + use super::*; + use litellm_secrets::{AccessMode, FailurePolicy, aws::AwsSecretsManagerV2}; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + + fn state(server: &MockServer, settings: KeyManagementSettings) -> SecretManagerState { + let endpoint = server.uri(); + let environment = Arc::new(move |name: &str| match name { + "AWS_REGION_NAME" => Some("us-east-1".into()), + "AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()), + "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()), + _ => None, + }); + let manager = + AwsSecretsManagerV2::load_aws_secret_manager(Some(true), settings.clone(), environment) + .unwrap() + .unwrap(); + SecretManagerState::new(SecretManager::AwsSecretsManagerV2(manager), settings) + } + + #[rstest::rstest] + #[case::missing(400, serde_json::json!({"__type":"ResourceNotFoundException"}), false)] + #[case::denied(400, serde_json::json!({"__type":"AccessDeniedException"}), true)] + #[case::malformed(200, serde_json::json!({}), true)] + #[tokio::test] + async fn failure_policy_preserves_errors_and_fallback_precedence( + #[case] status: u16, + #[case] body: serde_json::Value, + #[case] fails: bool, + #[values(FailurePolicy::Propagate, FailurePolicy::EnvironmentFallback)] + policy: FailurePolicy, + #[values(None, Some("environment"))] environment: Option<&'static str>, + #[values(None, Some("default"))] default: Option<&str>, + ) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(status).set_body_json(body)) + .expect(1) + .mount(&server) + .await; + let resolver = SecretResolver::new( + Arc::new(state(&server, KeyManagementSettings::default())), + Arc::new(move |_: &str| environment.map(str::to_owned)), + OidcResolver::default(), + ) + .with_failure_policy(policy); + let result = resolver + .get_secret_str("KEY", default.map(SecretValue::new)) + .await; + let fallback = environment.or(default); + if fails && (policy == FailurePolicy::Propagate || fallback.is_none()) { + assert!(matches!(result, Err(Error::Aws(_)))); + } else { + assert_eq!(result.unwrap().as_ref().map(SecretValue::expose), fallback); + } + } + + #[rstest::rstest] + #[case::boolean(serde_json::json!(false))] + #[case::object(serde_json::json!({"key":1}))] + #[case::null(serde_json::Value::Null)] + #[case::string(serde_json::json!("true"))] + #[tokio::test] + async fn typed_values_survive_resolution_and_accessors_reject_wrong_types( + #[case] value: serde_json::Value, + ) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json( + serde_json::json!({"SecretString":serde_json::json!({"KEY":value}).to_string()}), + )) + .expect(3) + .mount(&server) + .await; + let settings = KeyManagementSettings { + primary_secret_name: Some("primary".into()), + ..Default::default() + }; + let resolver = SecretResolver::new( + Arc::new(state(&server, settings)), + Arc::new(|_: &str| Some("fallback".into())), + OidcResolver::default(), + ); + assert_eq!( + resolver + .get_secret("KEY", Some(Secret::Bool(true))) + .await + .unwrap(), + Some(Secret::from_json(value.clone())) + ); + match &value { + serde_json::Value::String(text) => assert_eq!( + resolver + .get_secret_str("KEY", None) + .await + .unwrap() + .unwrap() + .expose(), + text + ), + _ => assert!(matches!( + resolver.get_secret_str("KEY", None).await, + Err(Error::TypeMismatch { expected: "string" }) + )), + } + match value { + serde_json::Value::Bool(boolean) => assert_eq!( + resolver.get_secret_bool("KEY", None).await.unwrap(), + Some(boolean) + ), + serde_json::Value::String(_) => assert_eq!( + resolver.get_secret_bool("KEY", None).await.unwrap(), + Some(true) + ), + _ => assert!(matches!( + resolver.get_secret_bool("KEY", None).await, + Err(Error::TypeMismatch { + expected: "boolean" + }) + )), + } + } + + #[rstest::rstest] + #[tokio::test] + async fn gating_prediction_matches_actual_lookup( + #[values(AccessMode::ReadOnly, AccessMode::WriteOnly, AccessMode::ReadAndWrite)] + access_mode: AccessMode, + #[values(None, Some(vec![]), Some(vec!["KEY".into()]))] hosted_keys: Option>, + #[values("os.environ/KEY", "os.environ/oidc/env/KEY")] name: &str, + ) { + let server = MockServer::start().await; + let expected = name == "os.environ/KEY" + && access_mode.readable() + && hosted_keys + .as_ref() + .is_none_or(|keys| keys.iter().any(|key| key == "KEY")); + Mock::given(method("POST")) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"SecretString":"remote"})), + ) + .expect(u64::from(expected)) + .mount(&server) + .await; + let state = state( + &server, + KeyManagementSettings { + access_mode, + hosted_keys, + ..Default::default() + }, + ); + assert!(state.backend().is_some()); + assert_eq!(state.settings().unwrap().access_mode, access_mode); + assert_eq!(secret_manager_would_be_consulted(&state, name), expected); + let resolver = SecretResolver::new( + Arc::new(state), + Arc::new(|_: &str| Some("environment".into())), + OidcResolver::default(), + ); + assert_eq!( + resolver + .get_secret_str(name, None) + .await + .unwrap() + .unwrap() + .expose(), + if expected { "remote" } else { "environment" } + ); + } +} + +#[cfg(feature = "google")] +#[rstest::rstest] +#[case::missing(404)] +#[case::failure(503)] +#[tokio::test] +async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16) { + use litellm_secrets::{FailurePolicy, google::GoogleSecretManager}; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + let server = MockServer::start().await; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(status)) + .expect(2) + .mount(&server) + .await; + let environment: Arc = + Arc::new(|name: &str| match name { + "VERTEX_AI_API_KEY" => Some("token".into()), + "KEY" => Some("environment".into()), + _ => None, + }); + let manager = GoogleSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + "project".into(), + environment.clone(), + None, + false, + ) + .unwrap(); + let state = SecretManagerState::new( + SecretManager::GoogleSecretManager(manager), + KeyManagementSettings::default(), + ); + let resolver = SecretResolver::new(Arc::new(state), environment, OidcResolver::default()); + let result = resolver.get_secret_str("KEY", None).await; + if status == 404 { + assert_eq!(result.unwrap().unwrap().expose(), "environment"); + } else { + assert!( + matches!(result, Err(Error::Google(litellm_secrets::google::Error::Status(actual))) if actual == status) + ); + } + assert_eq!( + resolver + .with_failure_policy(FailurePolicy::EnvironmentFallback) + .get_secret_str("KEY", None) + .await + .unwrap() + .unwrap() + .expose(), + "environment" + ); +} diff --git a/litellm-rust/crates/token-counter-fast/Cargo.toml b/litellm-rust/crates/token-counter-fast/Cargo.toml new file mode 100644 index 00000000000..5127dc17f22 --- /dev/null +++ b/litellm-rust/crates/token-counter-fast/Cargo.toml @@ -0,0 +1,19 @@ +[package] +name = "litellm-token-counter-fast" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +base64.workspace = true +rustc-hash = "2.1.3" +thiserror.workspace = true +tokenizers.workspace = true +unicode-normalization-alignments = "0.1.12" + +[dev-dependencies] +rand.workspace = true +rstest.workspace = true +serde.workspace = true +serde_json.workspace = true diff --git a/litellm-rust/crates/token-counter/src/byte_level.rs b/litellm-rust/crates/token-counter-fast/src/byte_level.rs similarity index 97% rename from litellm-rust/crates/token-counter/src/byte_level.rs rename to litellm-rust/crates/token-counter-fast/src/byte_level.rs index ec6134a252e..6fc7d9146b9 100644 --- a/litellm-rust/crates/token-counter/src/byte_level.rs +++ b/litellm-rust/crates/token-counter-fast/src/byte_level.rs @@ -473,13 +473,13 @@ mod tests { _ => unreachable!(), } assert!(ByteLevelCounter::detect(&anthropic_tokenizer).is_none()); - let counter = crate::TokenCounter::from_json( + let counter = crate::FastTokenizer::from_json( &anthropic_tokenizer.to_string(false).expect("serialize"), ) .expect("load"); for text in ["", "Hello WORLD! AB fi Ⅳ", " stop"] { assert_eq!( - counter.count_text(text).expect("count"), + counter.count_tokens(text).expect("count"), reference_count(&anthropic_tokenizer, text) ); } @@ -545,7 +545,7 @@ mod tests { .rstrip(rstrip)]) .expect("add token"); let fast = ByteLevelCounter::detect(&anthropic_tokenizer).expect("supported"); - let counter = crate::TokenCounter::from_json( + let counter = crate::FastTokenizer::from_json( &anthropic_tokenizer.to_string(false).expect("serialize"), ) .expect("load"); @@ -557,7 +557,7 @@ mod tests { ] { assert_eq!(fast.count(&anthropic_tokenizer, text), None); assert_eq!( - counter.count_text(text).expect("count"), + counter.count_tokens(text).expect("count"), reference_count(&anthropic_tokenizer, text) ); } @@ -571,17 +571,17 @@ mod tests { assert_eq!(fast.count(&tokenizer, "hello"), None); assert!(tokenizer.encode_fast("hello", true).is_err()); let counter = - crate::TokenCounter::from_json(&tokenizer.to_string(false).expect("serialize")) + crate::FastTokenizer::from_json(&tokenizer.to_string(false).expect("serialize")) .expect("load"); assert!(matches!( - counter.count_text("hello"), + counter.count_tokens("hello"), Err(crate::Error::Encode(_)) )); } #[rstest] fn shared_counter_matches_encoder_across_threads(anthropic_tokenizer: Tokenizer) { - let counter = crate::TokenCounter::from_json( + let counter = crate::FastTokenizer::from_json( &anthropic_tokenizer.to_string(false).expect("serialize"), ) .expect("load"); @@ -598,7 +598,7 @@ mod tests { scope.spawn(move || { for _ in 0..100 { for (text, count) in inputs.iter().zip(expected) { - assert_eq!(counter.count_text(text).expect("count"), count); + assert_eq!(counter.count_tokens(text).expect("count"), count); } } }); @@ -614,10 +614,10 @@ mod tests { let fast = ByteLevelCounter::detect(&anthropic_tokenizer).expect("supported"); assert_eq!(reference_count(&anthropic_tokenizer, "ABCD EFGH"), 1); assert_eq!(fast.count(&anthropic_tokenizer, "ABCD EFGH"), None); - let counter = crate::TokenCounter::from_json( + let counter = crate::FastTokenizer::from_json( &anthropic_tokenizer.to_string(false).expect("serialize"), ) .expect("load"); - assert_eq!(counter.count_text("ABCD EFGH").expect("count"), 1); + assert_eq!(counter.count_tokens("ABCD EFGH").expect("count"), 1); } } diff --git a/litellm-rust/crates/token-counter/src/cl100k.rs b/litellm-rust/crates/token-counter-fast/src/cl100k.rs similarity index 100% rename from litellm-rust/crates/token-counter/src/cl100k.rs rename to litellm-rust/crates/token-counter-fast/src/cl100k.rs diff --git a/litellm-rust/crates/token-counter-fast/src/error.rs b/litellm-rust/crates/token-counter-fast/src/error.rs new file mode 100644 index 00000000000..e63ccf3ad39 --- /dev/null +++ b/litellm-rust/crates/token-counter-fast/src/error.rs @@ -0,0 +1,13 @@ +use thiserror::Error as ThisError; + +#[derive(Debug, ThisError)] +pub enum Error { + #[error("failed to load tokenizer: {0}")] + Load(#[source] tokenizers::Error), + #[error("failed to load tokenizer: tiktoken rank file: {0}")] + Ranks(String), + #[error("failed to load tokenizer: Unicode character classes are unavailable")] + UnicodeClasses, + #[error("tokenization failed: {0}")] + Encode(#[source] tokenizers::Error), +} diff --git a/litellm-rust/crates/token-counter-fast/src/lib.rs b/litellm-rust/crates/token-counter-fast/src/lib.rs new file mode 100644 index 00000000000..ce91af642ea --- /dev/null +++ b/litellm-rust/crates/token-counter-fast/src/lib.rs @@ -0,0 +1,70 @@ +#![forbid(unsafe_code)] + +mod byte_level; +mod cl100k; +mod error; +mod o200k; +mod scanner; +mod tiktoken; +mod unicode_classes; + +use byte_level::ByteLevelCounter; +use scanner::{SplitPattern, TiktokenCounter}; + +pub use error::Error; + +enum Encoder { + HuggingFace { + tokenizer: Box, + byte_level: Option, + }, + Tiktoken(TiktokenCounter), +} + +pub struct FastTokenizer(Encoder); + +impl FastTokenizer { + pub fn from_json(json: &str) -> Result { + let tokenizer = json.parse::().map_err(Error::Load)?; + let byte_level = ByteLevelCounter::detect(&tokenizer); + Ok(Self(Encoder::HuggingFace { + tokenizer: Box::new(tokenizer), + byte_level, + })) + } + + pub fn from_cl100k_ranks(ranks: &str) -> Result { + Self::from_ranks(SplitPattern::Cl100k, ranks) + } + + pub fn from_o200k_ranks(ranks: &str) -> Result { + Self::from_ranks(SplitPattern::O200k, ranks) + } + + fn from_ranks(split: SplitPattern, ranks: &str) -> Result { + TiktokenCounter::from_ranks(split, ranks) + .map(Encoder::Tiktoken) + .map(Self) + } + + pub fn count_tokens(&self, text: &str) -> Result { + match &self.0 { + Encoder::Tiktoken(counter) => Ok(counter.count(text)), + Encoder::HuggingFace { + tokenizer, + byte_level, + } => { + if let Some(count) = byte_level + .as_ref() + .and_then(|counter| counter.count(tokenizer, text)) + { + return Ok(count); + } + tokenizer + .encode_fast(text, true) + .map(|encoding| encoding.len()) + .map_err(Error::Encode) + } + } + } +} diff --git a/litellm-rust/crates/token-counter/src/o200k.rs b/litellm-rust/crates/token-counter-fast/src/o200k.rs similarity index 100% rename from litellm-rust/crates/token-counter/src/o200k.rs rename to litellm-rust/crates/token-counter-fast/src/o200k.rs diff --git a/litellm-rust/crates/token-counter/src/scanner.rs b/litellm-rust/crates/token-counter-fast/src/scanner.rs similarity index 100% rename from litellm-rust/crates/token-counter/src/scanner.rs rename to litellm-rust/crates/token-counter-fast/src/scanner.rs diff --git a/litellm-rust/crates/token-counter-fast/src/tiktoken.rs b/litellm-rust/crates/token-counter-fast/src/tiktoken.rs new file mode 100644 index 00000000000..16172b7a688 --- /dev/null +++ b/litellm-rust/crates/token-counter-fast/src/tiktoken.rs @@ -0,0 +1,240 @@ +//! tiktoken's byte-level BPE: a rank file of `base64(token) rank` lines and +//! the merge loop that turns one regex piece into tokens. The merge order is +//! tiktoken's (lowest rank first, leftmost pair on ties) so the token count is +//! identical, but pairs are tracked in a heap so a long piece costs +//! `O(n log n)` instead of tiktoken's `O(n^2)`. + +use std::cmp::Reverse; +use std::collections::BinaryHeap; + +use base64::Engine; +use base64::engine::general_purpose::STANDARD; +use rustc_hash::FxHashMap; + +use crate::Error; + +type Rank = u32; + +const NO_RANK: Rank = Rank::MAX; +const END: usize = usize::MAX; + +pub(super) struct MergeRanks(FxHashMap, Rank>); + +impl MergeRanks { + pub(super) fn parse(text: &str) -> Result { + let ranks = text + .lines() + .filter(|line| !line.is_empty()) + .map(parse_line) + .collect::, _>>()?; + if let Some(byte) = (0..=u8::MAX).find(|byte| !ranks.contains_key(&[*byte][..])) { + return Err(Error::Ranks(format!("byte 0x{byte:02X} has no token"))); + } + Ok(Self(ranks)) + } + + fn rank(&self, bytes: &[u8]) -> Rank { + self.0.get(bytes).copied().unwrap_or(NO_RANK) + } + + /// Token count of one regex piece, as `encode_ordinary` would produce. + pub(super) fn count_piece(&self, piece: &[u8], scratch: &mut MergeScratch) -> usize { + if piece.len() < 2 || self.0.contains_key(piece) { + return 1; + } + scratch.reset(piece.len()); + for start in 0..piece.len() - 1 { + scratch.set_rank(start, self.rank(&piece[start..start + 2])); + } + let mut parts = piece.len(); + while let Some(Reverse((rank, start))) = scratch.heap.pop() { + if scratch.next[start] == END || scratch.rank[start] != rank { + continue; + } + let merged = scratch.next[start]; + let after = scratch.next[merged]; + scratch.next[merged] = END; + scratch.next[start] = after; + parts -= 1; + if after < piece.len() { + scratch.prev[after] = start; + scratch.set_rank(start, self.rank(&piece[start..scratch.end(after)])); + } else { + scratch.rank[start] = NO_RANK; + } + let before = scratch.prev[start]; + if before != END { + scratch.set_rank(before, self.rank(&piece[before..scratch.end(start)])); + } + } + parts + } +} + +fn parse_line(line: &str) -> Result<(Box<[u8]>, Rank), Error> { + let (token, rank) = line + .split_once(' ') + .ok_or_else(|| Error::Ranks(format!("line without a rank: {line:?}")))?; + let bytes = STANDARD + .decode(token) + .map_err(|error| Error::Ranks(format!("token is not base64: {error}")))?; + let rank = rank + .parse() + .map_err(|error| Error::Ranks(format!("rank is not an integer: {error}")))?; + if rank == NO_RANK { + return Err(Error::Ranks(format!("rank {rank} is reserved"))); + } + Ok((bytes.into_boxed_slice(), rank)) +} + +/// Buffers reused across the pieces of one text. Parts are addressed by the +/// byte offset they start at, which also gives the leftmost-pair tie break. +#[derive(Default)] +pub(super) struct MergeScratch { + next: Vec, + prev: Vec, + rank: Vec, + heap: BinaryHeap>, +} + +impl MergeScratch { + fn reset(&mut self, len: usize) { + self.next.clear(); + self.next.extend(1..=len); + self.prev.clear(); + self.prev.push(END); + self.prev.extend(0..len - 1); + self.rank.clear(); + self.rank.resize(len, NO_RANK); + self.heap.clear(); + } + + fn end(&self, start: usize) -> usize { + self.next[start] + } + + fn set_rank(&mut self, start: usize, rank: Rank) { + self.rank[start] = rank; + if rank != NO_RANK { + self.heap.push(Reverse((rank, start))); + } + } +} + +#[cfg(test)] +mod tests { + use rand::rngs::StdRng; + use rand::{Rng, SeedableRng}; + + use super::*; + + fn ranks() -> MergeRanks { + let path = concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../litellm/litellm_core_utils/tokenizers/9b5ad71b2ce5302211f9c61530b329a4922fc6a4" + ); + MergeRanks::parse(&std::fs::read_to_string(path).expect("cl100k rank file is in the repo")) + .expect("rank file parses") + } + + /// tiktoken's `_byte_pair_merge`, transcribed, as the reference. + fn reference_count(ranks: &MergeRanks, piece: &[u8]) -> usize { + if piece.len() < 2 || ranks.0.contains_key(piece) { + return 1; + } + let mut parts: Vec<(usize, Rank)> = (0..piece.len() - 1) + .map(|index| (index, ranks.rank(&piece[index..index + 2]))) + .chain([(piece.len() - 1, NO_RANK), (piece.len(), NO_RANK)]) + .collect(); + let get_rank = |parts: &[(usize, Rank)], index: usize| { + if index + 3 < parts.len() { + ranks.rank(&piece[parts[index].0..parts[index + 3].0]) + } else { + NO_RANK + } + }; + loop { + let Some(index) = parts[..parts.len() - 1] + .iter() + .enumerate() + .filter(|(_, (_, rank))| *rank != NO_RANK) + .min_by_key(|(index, (_, rank))| (*rank, *index)) + .map(|(index, _)| index) + else { + return parts.len() - 1; + }; + if index > 0 { + parts[index - 1].1 = get_rank(&parts, index - 1); + } + parts[index].1 = get_rank(&parts, index); + parts.remove(index + 1); + } + } + + #[test] + fn every_byte_is_a_token() { + let ranks = ranks(); + assert_eq!(ranks.0.len(), 100_256); + assert!((0..=u8::MAX).all(|byte| ranks.rank(&[byte]) != NO_RANK)); + } + + #[test] + fn heap_merge_matches_tiktokens_merge_loop() { + let ranks = ranks(); + let mut scratch = MergeScratch::default(); + let mut rng = StdRng::seed_from_u64(99); + let alphabet = b" abcdeorstn.,'\n\xc3\xa9\xe2\x82\xac0123"; + for _ in 0..20_000 { + let piece: Vec = (0..rng.gen_range(1..24)) + .map(|_| alphabet[rng.gen_range(0..alphabet.len())]) + .collect(); + assert_eq!( + ranks.count_piece(&piece, &mut scratch), + reference_count(&ranks, &piece), + "piece {:?}", + String::from_utf8_lossy(&piece) + ); + } + } + + #[test] + fn long_repeated_runs_cost_close_to_linear() { + let ranks = ranks(); + let mut scratch = MergeScratch::default(); + let mut time = |len: usize| { + let piece = vec![b' '; len]; + let started = std::time::Instant::now(); + assert!(ranks.count_piece(&piece, &mut scratch) > 0); + started.elapsed() + }; + let small = (0..3).map(|_| time(1 << 14)).min().unwrap(); + let large = time(1 << 18); + assert!( + large < small * 64, + "{small:?} for 2^14 bytes, {large:?} for 2^18" + ); + } + + #[test] + fn reserved_merge_rank_is_rejected() { + let bytes = (0..=u8::MAX) + .map(|byte| format!("{} {byte}\n", STANDARD.encode([byte]))) + .collect::(); + let rank_file = format!("{bytes}{} {NO_RANK}\n", STANDARD.encode(b"ab")); + assert!(matches!( + MergeRanks::parse(&rank_file), + Err(Error::Ranks(_)) + )); + let valid_rank_file = format!("{bytes}{} {}\n", STANDARD.encode(b"ab"), NO_RANK - 1); + let ranks = MergeRanks::parse(&valid_rank_file).unwrap(); + assert_eq!(ranks.count_piece(b"aab", &mut MergeScratch::default()), 2); + } + + #[test] + fn malformed_rank_files_are_rejected() { + assert!(MergeRanks::parse("IQ==").is_err()); + assert!(MergeRanks::parse("IQ== x").is_err()); + assert!(MergeRanks::parse("!!! 1").is_err()); + assert!(MergeRanks::parse("IQ== 1").is_err()); + } +} diff --git a/litellm-rust/crates/token-counter/src/unicode_classes.rs b/litellm-rust/crates/token-counter-fast/src/unicode_classes.rs similarity index 100% rename from litellm-rust/crates/token-counter/src/unicode_classes.rs rename to litellm-rust/crates/token-counter-fast/src/unicode_classes.rs diff --git a/litellm-rust/crates/token-counter/tests/fixtures/cl100k/requests.jsonl b/litellm-rust/crates/token-counter-fast/tests/fixtures/cl100k/requests.jsonl similarity index 100% rename from litellm-rust/crates/token-counter/tests/fixtures/cl100k/requests.jsonl rename to litellm-rust/crates/token-counter-fast/tests/fixtures/cl100k/requests.jsonl diff --git a/litellm-rust/crates/token-counter/tests/fixtures/cl100k/texts.jsonl b/litellm-rust/crates/token-counter-fast/tests/fixtures/cl100k/texts.jsonl similarity index 100% rename from litellm-rust/crates/token-counter/tests/fixtures/cl100k/texts.jsonl rename to litellm-rust/crates/token-counter-fast/tests/fixtures/cl100k/texts.jsonl diff --git a/litellm-rust/crates/token-counter/tests/fixtures/generate.py b/litellm-rust/crates/token-counter-fast/tests/fixtures/generate.py similarity index 98% rename from litellm-rust/crates/token-counter/tests/fixtures/generate.py rename to litellm-rust/crates/token-counter-fast/tests/fixtures/generate.py index 1bfbdf00218..2ca853cbd62 100644 --- a/litellm-rust/crates/token-counter/tests/fixtures/generate.py +++ b/litellm-rust/crates/token-counter-fast/tests/fixtures/generate.py @@ -2,8 +2,8 @@ Run from the repository root with the project environment, once per encoding: - uv run --no-sync python litellm-rust/crates/token-counter/tests/fixtures/generate.py cl100k_base - uv run --no-sync python litellm-rust/crates/token-counter/tests/fixtures/generate.py o200k_base + uv run --no-sync python litellm-rust/crates/token-counter-fast/tests/fixtures/generate.py cl100k_base + uv run --no-sync python litellm-rust/crates/token-counter-fast/tests/fixtures/generate.py o200k_base `/texts.jsonl` holds `{"text", "tokens", "pieces"}` lines: `tokens` counted with `tiktoken.get_encoding(name).encode(text, disallowed_special=())`, diff --git a/litellm-rust/crates/token-counter/tests/fixtures/o200k/requests.jsonl b/litellm-rust/crates/token-counter-fast/tests/fixtures/o200k/requests.jsonl similarity index 100% rename from litellm-rust/crates/token-counter/tests/fixtures/o200k/requests.jsonl rename to litellm-rust/crates/token-counter-fast/tests/fixtures/o200k/requests.jsonl diff --git a/litellm-rust/crates/token-counter/tests/fixtures/o200k/texts.jsonl b/litellm-rust/crates/token-counter-fast/tests/fixtures/o200k/texts.jsonl similarity index 100% rename from litellm-rust/crates/token-counter/tests/fixtures/o200k/texts.jsonl rename to litellm-rust/crates/token-counter-fast/tests/fixtures/o200k/texts.jsonl diff --git a/litellm-rust/crates/token-counter-huggingface/Cargo.toml b/litellm-rust/crates/token-counter-huggingface/Cargo.toml new file mode 100644 index 00000000000..6d8cb85e524 --- /dev/null +++ b/litellm-rust/crates/token-counter-huggingface/Cargo.toml @@ -0,0 +1,10 @@ +[package] +name = "litellm-token-counter-huggingface" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +thiserror.workspace = true +tokenizers.workspace = true diff --git a/litellm-rust/crates/token-counter-huggingface/src/error.rs b/litellm-rust/crates/token-counter-huggingface/src/error.rs new file mode 100644 index 00000000000..adc4551886f --- /dev/null +++ b/litellm-rust/crates/token-counter-huggingface/src/error.rs @@ -0,0 +1,9 @@ +use thiserror::Error as ThisError; + +#[derive(Debug, ThisError)] +pub enum Error { + #[error("failed to load tokenizer: {0}")] + Load(#[source] tokenizers::Error), + #[error("tokenization failed: {0}")] + Encode(#[source] tokenizers::Error), +} diff --git a/litellm-rust/crates/token-counter-huggingface/src/lib.rs b/litellm-rust/crates/token-counter-huggingface/src/lib.rs new file mode 100644 index 00000000000..8e05c2cca46 --- /dev/null +++ b/litellm-rust/crates/token-counter-huggingface/src/lib.rs @@ -0,0 +1,23 @@ +#![forbid(unsafe_code)] + +mod error; + +pub use error::Error; + +pub struct HuggingFaceTokenizer(Box); + +impl HuggingFaceTokenizer { + pub fn from_json(json: &str) -> Result { + json.parse::() + .map(Box::new) + .map(Self) + .map_err(Error::Load) + } + + pub fn count_tokens(&self, text: &str) -> Result { + self.0 + .encode_fast(text, true) + .map(|encoding| encoding.len()) + .map_err(Error::Encode) + } +} diff --git a/litellm-rust/crates/token-counter-tiktoken/Cargo.toml b/litellm-rust/crates/token-counter-tiktoken/Cargo.toml new file mode 100644 index 00000000000..494a9233e69 --- /dev/null +++ b/litellm-rust/crates/token-counter-tiktoken/Cargo.toml @@ -0,0 +1,10 @@ +[package] +name = "litellm-token-counter-tiktoken" +version = "0.1.0" +edition.workspace = true +license.workspace = true +repository.workspace = true + +[dependencies] +thiserror.workspace = true +tiktoken-rs.workspace = true diff --git a/litellm-rust/crates/token-counter-tiktoken/src/error.rs b/litellm-rust/crates/token-counter-tiktoken/src/error.rs new file mode 100644 index 00000000000..e28cbbfb620 --- /dev/null +++ b/litellm-rust/crates/token-counter-tiktoken/src/error.rs @@ -0,0 +1,5 @@ +use thiserror::Error as ThisError; + +#[derive(Debug, ThisError)] +#[error("unsupported tokenizer: {0}")] +pub struct UnsupportedTokenizer(pub String); diff --git a/litellm-rust/crates/token-counter-tiktoken/src/lib.rs b/litellm-rust/crates/token-counter-tiktoken/src/lib.rs new file mode 100644 index 00000000000..ecdb3946eee --- /dev/null +++ b/litellm-rust/crates/token-counter-tiktoken/src/lib.rs @@ -0,0 +1,70 @@ +#![forbid(unsafe_code)] + +mod error; + +pub use error::UnsupportedTokenizer; + +pub struct TiktokenTokenizer(&'static tiktoken_rs::CoreBPE); + +impl TiktokenTokenizer { + pub fn from_name(name: &str) -> Result { + let tokenizer = match name { + "cl100k_base" => tiktoken_rs::cl100k_base_singleton(), + "o200k_base" => tiktoken_rs::o200k_base_singleton(), + "o200k_harmony" => tiktoken_rs::o200k_harmony_singleton(), + "p50k_base" => tiktoken_rs::p50k_base_singleton(), + "p50k_edit" => tiktoken_rs::p50k_edit_singleton(), + "r50k_base" | "gpt2" => tiktoken_rs::r50k_base_singleton(), + _ => return Err(UnsupportedTokenizer(name.to_owned())), + }; + Ok(Self(tokenizer)) + } + + pub fn count_tokens(&self, text: &str) -> usize { + self.0.count_ordinary(text) + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn named_encodings_match_their_reference_counts() { + let encodings = [ + ("cl100k_base", tiktoken_rs::cl100k_base_singleton()), + ("o200k_base", tiktoken_rs::o200k_base_singleton()), + ("o200k_harmony", tiktoken_rs::o200k_harmony_singleton()), + ("p50k_base", tiktoken_rs::p50k_base_singleton()), + ("p50k_edit", tiktoken_rs::p50k_edit_singleton()), + ("r50k_base", tiktoken_rs::r50k_base_singleton()), + ("gpt2", tiktoken_rs::r50k_base_singleton()), + ]; + let texts = [ + "", + "Hello, how are you today?", + "é e\u{301} 漢字 ع ३ 🙂 AfiⅣ", + " def function():\n return 123456789\r\n", + "<|endoftext|><|fim_prefix|><|start|>assistant<|message|>", + ]; + for (name, reference) in encodings { + let counter = TiktokenTokenizer::from_name(name).unwrap(); + for text in texts { + assert_eq!( + counter.count_tokens(text), + reference.encode_ordinary(text).len(), + "{name}: {text:?}", + ); + } + } + } + + #[test] + fn unsupported_encoding_preserves_its_name() { + let Err(UnsupportedTokenizer(name)) = TiktokenTokenizer::from_name("unknown-encoding") + else { + panic!("unknown encoding must be rejected"); + }; + assert_eq!(name, "unknown-encoding"); + } +} diff --git a/litellm-rust/crates/token-counter/Cargo.toml b/litellm-rust/crates/token-counter/Cargo.toml index d0369631682..67e5bd60537 100644 --- a/litellm-rust/crates/token-counter/Cargo.toml +++ b/litellm-rust/crates/token-counter/Cargo.toml @@ -5,26 +5,34 @@ edition.workspace = true license.workspace = true repository.workspace = true +[features] +default = ["fast", "huggingface", "tiktoken"] +fast = ["dep:litellm-token-counter-fast"] +huggingface = ["dep:litellm-token-counter-huggingface"] +tiktoken = ["dep:litellm-token-counter-tiktoken"] + [dependencies] -base64.workspace = true indexmap = { version = "2.14.0", features = ["serde"] } itoa = "1.0" -rustc-hash = "2.1.3" +litellm-token-counter-fast = { workspace = true, optional = true } +litellm-token-counter-huggingface = { workspace = true, optional = true } +litellm-token-counter-tiktoken = { workspace = true, optional = true } serde.workspace = true serde_json.workspace = true thiserror.workspace = true -tokenizers = { version = "0.23.1", default-features = false, features = ["onig"] } -unicode-normalization-alignments = "0.1.12" [dev-dependencies] criterion.workspace = true rand.workspace = true rstest.workspace = true +tokenizers.workspace = true [[bench]] name = "token_counter" harness = false +required-features = ["fast"] [[bench]] name = "allocations" harness = false +required-features = ["fast"] diff --git a/litellm-rust/crates/token-counter/README.md b/litellm-rust/crates/token-counter/README.md new file mode 100644 index 00000000000..a6b2b50aac0 --- /dev/null +++ b/litellm-rust/crates/token-counter/README.md @@ -0,0 +1,23 @@ +# Token counting + +`Tokenizer` is the text-counting interface. `TokenCounter` applies LiteLLM request, message, and tool accounting using any implementation of that interface + +The `fast` feature provides `fast::FastTokenizer` from `litellm-token-counter-fast`. `TokenCounter::from_json_fast` uses this implementation + +The `huggingface` feature provides `huggingface::HuggingFaceTokenizer` through the upstream `tokenizers` library. `TokenCounter::from_json` uses this implementation + +The `tiktoken` feature provides `tiktoken::TiktokenTokenizer` through `tiktoken-rs`. Select an encoding with `TokenCounter::from_tiktoken`. The supported names are `cl100k_base`, `o200k_base`, `o200k_harmony`, `p50k_base`, `p50k_edit`, `r50k_base`, and `gpt2` + +All three backends are enabled by default. The Python extension builds with `fast` only, which keeps the wheel at the size it had before the split. With `default-features = false`, callers can supply their own `Tokenizer` to `TokenCounter::new` without compiling a built-in backend + +Budget checks, cost calculation, and the `max_tokens` adjustment policy belong to `litellm-core-utils`. The counter does not own prices, budgets, or request limits + +Run the feature matrix with: + +```sh +cargo test -p litellm-token-counter +cargo test -p litellm-token-counter --no-default-features +cargo test -p litellm-token-counter --no-default-features --features fast +cargo test -p litellm-token-counter --no-default-features --features huggingface +cargo test -p litellm-token-counter --no-default-features --features tiktoken +``` diff --git a/litellm-rust/crates/token-counter/benches/allocations.rs b/litellm-rust/crates/token-counter/benches/allocations.rs index 343f815749d..7aebceb6e9c 100644 --- a/litellm-rust/crates/token-counter/benches/allocations.rs +++ b/litellm-rust/crates/token-counter/benches/allocations.rs @@ -87,7 +87,7 @@ fn main() { }, ); - let counter = TokenCounter::from_json(TOKENIZER_JSON).expect("tokenizer loads"); + let counter = TokenCounter::from_json_fast(TOKENIZER_JSON).expect("tokenizer loads"); let object = CountableRequest::parse(OBJECT_BODY).expect("object request parses"); counter .count_request(&object) diff --git a/litellm-rust/crates/token-counter/benches/token_counter.rs b/litellm-rust/crates/token-counter/benches/token_counter.rs index a7c3177b88a..5e30cf6f77e 100644 --- a/litellm-rust/crates/token-counter/benches/token_counter.rs +++ b/litellm-rust/crates/token-counter/benches/token_counter.rs @@ -42,7 +42,7 @@ fn inputs(tokenizer: &Tokenizer) -> Vec<(&'static str, String)> { } fn token_counter(c: &mut Criterion) { - let counter = TokenCounter::from_json(TOKENIZER_JSON).expect("token counter should load"); + let counter = TokenCounter::from_json_fast(TOKENIZER_JSON).expect("token counter should load"); let tokenizer = TOKENIZER_JSON .parse::() .expect("reference tokenizer should load"); diff --git a/litellm-rust/crates/token-counter/src/counter.rs b/litellm-rust/crates/token-counter/src/counter.rs index 7eedc449dd1..ce08e225be4 100644 --- a/litellm-rust/crates/token-counter/src/counter.rs +++ b/litellm-rust/crates/token-counter/src/counter.rs @@ -1,9 +1,7 @@ use serde::Serialize; use crate::Error; -use crate::byte_level::ByteLevelCounter; use crate::python_json; -use crate::scanner::{SplitPattern, TiktokenCounter}; use crate::tools::format_function_definitions; use crate::types::{ ContentBlock, ContentItem, CountableRequest, Message, MessageContent, TextValue, ToolChoice, @@ -24,73 +22,22 @@ pub struct InputTokenCount { pub input_tokens: usize, } -enum Encoder { - HuggingFace { - tokenizer: Box, - byte_level: Option, - }, - Tiktoken(TiktokenCounter), -} - /// A loaded tokenizer plus the message accounting Python applies on top of /// it. Encoding is CPU-bound and synchronous; hosts run it off their event /// loop. pub struct TokenCounter { - encoder: Encoder, + encoder: Box, } impl TokenCounter { - /// Load a HuggingFace `tokenizer.json` document. The host reads the file. - pub fn from_json(tokenizer_json: &str) -> Result { - let tokenizer = tokenizer_json - .parse::() - .map_err(Error::Load)?; - let byte_level = ByteLevelCounter::detect(&tokenizer); - Ok(Self { - encoder: Encoder::HuggingFace { - tokenizer: Box::new(tokenizer), - byte_level, - }, - }) - } - - /// Load tiktoken's `cl100k_base` rank file (`base64(token) rank` lines). - /// The host reads the file. - pub fn from_cl100k_ranks(rank_file: &str) -> Result { - Self::from_tiktoken_ranks(SplitPattern::Cl100k, rank_file) - } - - /// Load tiktoken's `o200k_base` rank file (`base64(token) rank` lines). - /// The host reads the file. - pub fn from_o200k_ranks(rank_file: &str) -> Result { - Self::from_tiktoken_ranks(SplitPattern::O200k, rank_file) - } - - fn from_tiktoken_ranks(split: SplitPattern, rank_file: &str) -> Result { - Ok(Self { - encoder: Encoder::Tiktoken(TiktokenCounter::from_ranks(split, rank_file)?), - }) + pub fn new(tokenizer: impl crate::Tokenizer + 'static) -> Self { + Self { + encoder: Box::new(tokenizer), + } } pub fn count_text(&self, text: &str) -> Result { - match &self.encoder { - Encoder::Tiktoken(counter) => Ok(counter.count(text)), - Encoder::HuggingFace { - tokenizer, - byte_level, - } => { - if let Some(count) = byte_level - .as_ref() - .and_then(|counter| counter.count(tokenizer, text)) - { - return Ok(count); - } - tokenizer - .encode_fast(text, true) - .map(|encoding| encoding.len()) - .map_err(Error::Encode) - } - } + self.encoder.count_tokens(text) } /// Mirrors the host's key precedence: `messages`, then `prompt`, then diff --git a/litellm-rust/crates/token-counter/src/error.rs b/litellm-rust/crates/token-counter/src/error.rs index 6b8668fe182..b05ce007e46 100644 --- a/litellm-rust/crates/token-counter/src/error.rs +++ b/litellm-rust/crates/token-counter/src/error.rs @@ -4,8 +4,10 @@ use thiserror::Error as ThisError; #[derive(Debug, ThisError)] pub enum Error { + #[error("unsupported tokenizer: {0}")] + UnsupportedTokenizer(String), #[error("failed to load tokenizer: {0}")] - Load(#[source] tokenizers::Error), + Load(#[source] Box), #[error("failed to load tokenizer: tiktoken rank file: {0}")] Ranks(String), #[error("failed to load tokenizer: Unicode character classes are unavailable")] @@ -29,7 +31,7 @@ pub enum Error { #[error("unsupported by the rust token counter: serialized text value is not UTF-8: {0}")] JsonUtf8(#[source] FromUtf8Error), #[error("tokenization failed: {0}")] - Encode(#[source] tokenizers::Error), + Encode(#[source] Box), #[error("token counting task failed: {0}")] Task(String), } diff --git a/litellm-rust/crates/token-counter/src/fast.rs b/litellm-rust/crates/token-counter/src/fast.rs new file mode 100644 index 00000000000..de3f86abd68 --- /dev/null +++ b/litellm-rust/crates/token-counter/src/fast.rs @@ -0,0 +1,41 @@ +use litellm_token_counter_fast::Error as BackendError; +pub use litellm_token_counter_fast::FastTokenizer; + +use crate::{Error, TokenCounter, Tokenizer}; + +impl TokenCounter { + pub fn from_json_fast(tokenizer_json: &str) -> Result { + FastTokenizer::from_json(tokenizer_json) + .map(Self::new) + .map_err(Error::from) + } + + pub fn from_cl100k_ranks(rank_file: &str) -> Result { + FastTokenizer::from_cl100k_ranks(rank_file) + .map(Self::new) + .map_err(Error::from) + } + + pub fn from_o200k_ranks(rank_file: &str) -> Result { + FastTokenizer::from_o200k_ranks(rank_file) + .map(Self::new) + .map_err(Error::from) + } +} + +impl Tokenizer for FastTokenizer { + fn count_tokens(&self, text: &str) -> Result { + FastTokenizer::count_tokens(self, text).map_err(Error::from) + } +} + +impl From for Error { + fn from(error: BackendError) -> Self { + match error { + BackendError::Load(source) => Self::Load(source), + BackendError::Ranks(message) => Self::Ranks(message), + BackendError::UnicodeClasses => Self::UnicodeClasses, + BackendError::Encode(source) => Self::Encode(source), + } + } +} diff --git a/litellm-rust/crates/token-counter/src/huggingface.rs b/litellm-rust/crates/token-counter/src/huggingface.rs new file mode 100644 index 00000000000..fb7683b373e --- /dev/null +++ b/litellm-rust/crates/token-counter/src/huggingface.rs @@ -0,0 +1,27 @@ +use litellm_token_counter_huggingface::Error as BackendError; +pub use litellm_token_counter_huggingface::HuggingFaceTokenizer; + +use crate::{Error, TokenCounter, Tokenizer}; + +impl TokenCounter { + pub fn from_json(tokenizer_json: &str) -> Result { + HuggingFaceTokenizer::from_json(tokenizer_json) + .map(Self::new) + .map_err(Error::from) + } +} + +impl Tokenizer for HuggingFaceTokenizer { + fn count_tokens(&self, text: &str) -> Result { + HuggingFaceTokenizer::count_tokens(self, text).map_err(Error::from) + } +} + +impl From for Error { + fn from(error: BackendError) -> Self { + match error { + BackendError::Load(source) => Self::Load(source), + BackendError::Encode(source) => Self::Encode(source), + } + } +} diff --git a/litellm-rust/crates/token-counter/src/lib.rs b/litellm-rust/crates/token-counter/src/lib.rs index fa0014e2bad..446c91049de 100644 --- a/litellm-rust/crates/token-counter/src/lib.rs +++ b/litellm-rust/crates/token-counter/src/lib.rs @@ -4,18 +4,21 @@ #![forbid(unsafe_code)] -mod byte_level; -mod cl100k; mod counter; mod error; -mod o200k; mod python_json; -mod scanner; -mod tiktoken; +mod tokenizer; mod tools; mod types; -mod unicode_classes; + +#[cfg(feature = "fast")] +pub mod fast; +#[cfg(feature = "huggingface")] +pub mod huggingface; +#[cfg(feature = "tiktoken")] +pub mod tiktoken; pub use counter::{InputTokenCount, TokenCounter}; pub use error::Error; +pub use tokenizer::Tokenizer; pub use types::CountableRequest; diff --git a/litellm-rust/crates/token-counter/src/tiktoken.rs b/litellm-rust/crates/token-counter/src/tiktoken.rs index 7a9e71ed587..07c1c9f5b73 100644 --- a/litellm-rust/crates/token-counter/src/tiktoken.rs +++ b/litellm-rust/crates/token-counter/src/tiktoken.rs @@ -1,222 +1,24 @@ -//! tiktoken's byte-level BPE: a rank file of `base64(token) rank` lines and -//! the merge loop that turns one regex piece into tokens. The merge order is -//! tiktoken's (lowest rank first, leftmost pair on ties) so the token count is -//! identical, but pairs are tracked in a heap so a long piece costs -//! `O(n log n)` instead of tiktoken's `O(n^2)`. +pub use litellm_token_counter_tiktoken::TiktokenTokenizer; +use litellm_token_counter_tiktoken::UnsupportedTokenizer; -use std::cmp::Reverse; -use std::collections::BinaryHeap; +use crate::{Error, TokenCounter, Tokenizer}; -use base64::Engine; -use base64::engine::general_purpose::STANDARD; -use rustc_hash::FxHashMap; - -use crate::Error; - -type Rank = u32; - -const NO_RANK: Rank = Rank::MAX; -const END: usize = usize::MAX; - -pub(super) struct MergeRanks(FxHashMap, Rank>); - -impl MergeRanks { - pub(super) fn parse(text: &str) -> Result { - let ranks = text - .lines() - .filter(|line| !line.is_empty()) - .map(parse_line) - .collect::, _>>()?; - if let Some(byte) = (0..=u8::MAX).find(|byte| !ranks.contains_key(&[*byte][..])) { - return Err(Error::Ranks(format!("byte 0x{byte:02X} has no token"))); - } - Ok(Self(ranks)) - } - - fn rank(&self, bytes: &[u8]) -> Rank { - self.0.get(bytes).copied().unwrap_or(NO_RANK) - } - - /// Token count of one regex piece, as `encode_ordinary` would produce. - pub(super) fn count_piece(&self, piece: &[u8], scratch: &mut MergeScratch) -> usize { - if piece.len() < 2 || self.0.contains_key(piece) { - return 1; - } - scratch.reset(piece.len()); - for start in 0..piece.len() - 1 { - scratch.set_rank(start, self.rank(&piece[start..start + 2])); - } - let mut parts = piece.len(); - while let Some(Reverse((rank, start))) = scratch.heap.pop() { - if scratch.next[start] == END || scratch.rank[start] != rank { - continue; - } - let merged = scratch.next[start]; - let after = scratch.next[merged]; - scratch.next[merged] = END; - scratch.next[start] = after; - parts -= 1; - if after < piece.len() { - scratch.prev[after] = start; - scratch.set_rank(start, self.rank(&piece[start..scratch.end(after)])); - } else { - scratch.rank[start] = NO_RANK; - } - let before = scratch.prev[start]; - if before != END { - scratch.set_rank(before, self.rank(&piece[before..scratch.end(start)])); - } - } - parts +impl TokenCounter { + pub fn from_tiktoken(encoding: &str) -> Result { + TiktokenTokenizer::from_name(encoding) + .map(Self::new) + .map_err(Error::from) } } -fn parse_line(line: &str) -> Result<(Box<[u8]>, Rank), Error> { - let (token, rank) = line - .split_once(' ') - .ok_or_else(|| Error::Ranks(format!("line without a rank: {line:?}")))?; - let bytes = STANDARD - .decode(token) - .map_err(|error| Error::Ranks(format!("token is not base64: {error}")))?; - let rank = rank - .parse() - .map_err(|error| Error::Ranks(format!("rank is not an integer: {error}")))?; - Ok((bytes.into_boxed_slice(), rank)) -} - -/// Buffers reused across the pieces of one text. Parts are addressed by the -/// byte offset they start at, which also gives the leftmost-pair tie break. -#[derive(Default)] -pub(super) struct MergeScratch { - next: Vec, - prev: Vec, - rank: Vec, - heap: BinaryHeap>, -} - -impl MergeScratch { - fn reset(&mut self, len: usize) { - self.next.clear(); - self.next.extend(1..=len); - self.prev.clear(); - self.prev.push(END); - self.prev.extend(0..len - 1); - self.rank.clear(); - self.rank.resize(len, NO_RANK); - self.heap.clear(); - } - - fn end(&self, start: usize) -> usize { - self.next[start] - } - - fn set_rank(&mut self, start: usize, rank: Rank) { - self.rank[start] = rank; - if rank != NO_RANK { - self.heap.push(Reverse((rank, start))); - } +impl Tokenizer for TiktokenTokenizer { + fn count_tokens(&self, text: &str) -> Result { + Ok(TiktokenTokenizer::count_tokens(self, text)) } } -#[cfg(test)] -mod tests { - use rand::rngs::StdRng; - use rand::{Rng, SeedableRng}; - - use super::*; - - fn ranks() -> MergeRanks { - let path = concat!( - env!("CARGO_MANIFEST_DIR"), - "/../../../litellm/litellm_core_utils/tokenizers/9b5ad71b2ce5302211f9c61530b329a4922fc6a4" - ); - MergeRanks::parse(&std::fs::read_to_string(path).expect("cl100k rank file is in the repo")) - .expect("rank file parses") - } - - /// tiktoken's `_byte_pair_merge`, transcribed, as the reference. - fn reference_count(ranks: &MergeRanks, piece: &[u8]) -> usize { - if piece.len() < 2 || ranks.0.contains_key(piece) { - return 1; - } - let mut parts: Vec<(usize, Rank)> = (0..piece.len() - 1) - .map(|index| (index, ranks.rank(&piece[index..index + 2]))) - .chain([(piece.len() - 1, NO_RANK), (piece.len(), NO_RANK)]) - .collect(); - let get_rank = |parts: &[(usize, Rank)], index: usize| { - if index + 3 < parts.len() { - ranks.rank(&piece[parts[index].0..parts[index + 3].0]) - } else { - NO_RANK - } - }; - loop { - let Some(index) = parts[..parts.len() - 1] - .iter() - .enumerate() - .filter(|(_, (_, rank))| *rank != NO_RANK) - .min_by_key(|(index, (_, rank))| (*rank, *index)) - .map(|(index, _)| index) - else { - return parts.len() - 1; - }; - if index > 0 { - parts[index - 1].1 = get_rank(&parts, index - 1); - } - parts[index].1 = get_rank(&parts, index); - parts.remove(index + 1); - } - } - - #[test] - fn every_byte_is_a_token() { - let ranks = ranks(); - assert_eq!(ranks.0.len(), 100_256); - assert!((0..=u8::MAX).all(|byte| ranks.rank(&[byte]) != NO_RANK)); - } - - #[test] - fn heap_merge_matches_tiktokens_merge_loop() { - let ranks = ranks(); - let mut scratch = MergeScratch::default(); - let mut rng = StdRng::seed_from_u64(99); - let alphabet = b" abcdeorstn.,'\n\xc3\xa9\xe2\x82\xac0123"; - for _ in 0..20_000 { - let piece: Vec = (0..rng.gen_range(1..24)) - .map(|_| alphabet[rng.gen_range(0..alphabet.len())]) - .collect(); - assert_eq!( - ranks.count_piece(&piece, &mut scratch), - reference_count(&ranks, &piece), - "piece {:?}", - String::from_utf8_lossy(&piece) - ); - } - } - - #[test] - fn long_repeated_runs_cost_close_to_linear() { - let ranks = ranks(); - let mut scratch = MergeScratch::default(); - let mut time = |len: usize| { - let piece = vec![b' '; len]; - let started = std::time::Instant::now(); - assert!(ranks.count_piece(&piece, &mut scratch) > 0); - started.elapsed() - }; - let small = (0..3).map(|_| time(1 << 14)).min().unwrap(); - let large = time(1 << 18); - assert!( - large < small * 64, - "{small:?} for 2^14 bytes, {large:?} for 2^18" - ); - } - - #[test] - fn malformed_rank_files_are_rejected() { - assert!(MergeRanks::parse("IQ==").is_err()); - assert!(MergeRanks::parse("IQ== x").is_err()); - assert!(MergeRanks::parse("!!! 1").is_err()); - assert!(MergeRanks::parse("IQ== 1").is_err()); +impl From for Error { + fn from(error: UnsupportedTokenizer) -> Self { + Self::UnsupportedTokenizer(error.0) } } diff --git a/litellm-rust/crates/token-counter/src/tokenizer.rs b/litellm-rust/crates/token-counter/src/tokenizer.rs new file mode 100644 index 00000000000..88c29c672a7 --- /dev/null +++ b/litellm-rust/crates/token-counter/src/tokenizer.rs @@ -0,0 +1,31 @@ +use crate::Error; + +pub trait Tokenizer: Send + Sync { + fn count_tokens(&self, text: &str) -> Result; +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::{CountableRequest, TokenCounter}; + + struct Characters; + + impl Tokenizer for Characters { + fn count_tokens(&self, text: &str) -> Result { + Ok(text.chars().count()) + } + } + + #[test] + fn request_accounting_works_with_an_injected_backend() { + let counter = TokenCounter::new(Characters); + let request = + CountableRequest::parse(br#"{"messages":[{"role":"user","content":"hello"}]}"#) + .unwrap(); + assert_eq!( + counter.count_request(&request).unwrap().input_tokens, + 3 + 4 + 5 + 3 + ); + } +} diff --git a/litellm-rust/crates/token-counter/tests/token_counter.rs b/litellm-rust/crates/token-counter/tests/token_counter.rs index 12c59768952..542bd4a1fc4 100644 --- a/litellm-rust/crates/token-counter/tests/token_counter.rs +++ b/litellm-rust/crates/token-counter/tests/token_counter.rs @@ -1,22 +1,30 @@ use rstest::rstest; -use serde::Deserialize; -use litellm_token_counter::{CountableRequest, Error, InputTokenCount, TokenCounter}; +#[cfg(any(feature = "fast", feature = "huggingface", feature = "tiktoken"))] +use litellm_token_counter::TokenCounter; +use litellm_token_counter::{CountableRequest, Error}; -/// Expected counts are pinned from `litellm.token_counter(model="claude-sonnet-4-5", ...)` -/// so this test also guards Python parity. -fn counter() -> TokenCounter { - let path = concat!( - env!("CARGO_MANIFEST_DIR"), - "/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json" - ); - let json = std::fs::read_to_string(path).expect("anthropic tokenizer json is in the repo"); - TokenCounter::from_json(&json).expect("anthropic tokenizer loads") -} +#[cfg(any(feature = "fast", feature = "huggingface"))] +mod json { + use super::*; + use litellm_token_counter::InputTokenCount; -const SIMPLE: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"Hello, how are you today?"}]}"#; + /// Expected counts are pinned from `litellm.token_counter(model="claude-sonnet-4-5", ...)` + /// so this test also guards Python parity. + type JsonLoader = fn(&str) -> Result; -const BLOCKS_AND_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5","messages":[ + fn counter(load: JsonLoader) -> TokenCounter { + let path = concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../litellm/litellm_core_utils/tokenizers/anthropic_tokenizer.json" + ); + let json = std::fs::read_to_string(path).expect("anthropic tokenizer json is in the repo"); + load(&json).expect("anthropic tokenizer loads") + } + + const SIMPLE: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"Hello, how are you today?"}]}"#; + + const BLOCKS_AND_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5","messages":[ {"role":"system","content":"You are a terse assistant."}, {"role":"user","name":"alice","content":[ {"type":"text","text":"Summarise this paragraph about ships and harbours."}, @@ -25,7 +33,7 @@ const BLOCKS_AND_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5","messages":[ {"type":"tool_reference","tool_name":"get_weather"}]}, {"role":"assistant","content":[{"type":"text","text":"Sure.","cache_control":{"type":"ephemeral"}}]}]}"#; -const TOOLS_OPENAI: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"weather?"}], + const TOOLS_OPENAI: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":"user","content":"weather?"}], "tools":[ {"type":"function","function":{"name":"get_weather","description":"Get weather","parameters":{ "type":"object", @@ -40,61 +48,188 @@ const TOOLS_OPENAI: &str = r#"{"model":"claude-sonnet-4-5","messages":[{"role":" {"type":"function","function":{"name":"noop"}}], "tool_choice":{"type":"function","function":{"name":"get_weather"}}}"#; -const TOOLS_ANTHROPIC_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5", + const TOOLS_ANTHROPIC_SYSTEM: &str = r#"{"model":"claude-sonnet-4-5", "messages":[{"role":"system","content":"sys"},{"role":"user","content":"weather?"}], "tools":[{"name":"get_weather","description":"Get weather","input_schema":{ "type":"object","properties":{"location":{"type":["string","null"]}},"required":["location"]}}], "tool_choice":"none"}"#; -const COMPLETIONS_PROMPT: &str = - r#"{"model":"claude-sonnet-4-5","prompt":"Write a haiku about ships."}"#; + const COMPLETIONS_PROMPT: &str = + r#"{"model":"claude-sonnet-4-5","prompt":"Write a haiku about ships."}"#; -const COMPLETIONS_PROMPT_LIST: &str = - r#"{"model":"claude-sonnet-4-5","prompt":["first prompt","second prompt"]}"#; + const COMPLETIONS_PROMPT_LIST: &str = + r#"{"model":"claude-sonnet-4-5","prompt":["first prompt","second prompt"]}"#; -const RESPONSES_INPUT: &str = r#"{"model":"claude-sonnet-4-5","input":[ + const RESPONSES_INPUT: &str = r#"{"model":"claude-sonnet-4-5","input":[ {"role":"user","content":[{"type":"input_text","text":"Summarise caf\u00e9 menus, na\u00efve \u2014 ok? \"quoted\"\n"}]}, {"role":"assistant","content":"Sure."}],"instructions":"be terse"}"#; -const EMBEDDINGS_TOKEN_IDS: &str = - r#"{"model":"claude-sonnet-4-5","input":[[101,2023,5],[7]],"encoding_format":"float"}"#; + const EMBEDDINGS_TOKEN_IDS: &str = + r#"{"model":"claude-sonnet-4-5","input":[[101,2023,5],[7]],"encoding_format":"float"}"#; -const RERANK: &str = r#"{"model":"claude-sonnet-4-5","query":"best harbour", + const RERANK: &str = r#"{"model":"claude-sonnet-4-5","query":"best harbour", "documents":["doc one",{"text":"doc two","title":"T","n":3,"ok":true,"none":null,"tags":["a","b"]}]}"#; -/// Expected counts are pinned from -/// `litellm.proxy.spend_tracking.budget_reservation._count_input_tokens(body, "claude-sonnet-4-5")`. -#[rstest] -#[case::text_only(SIMPLE, 14)] -#[case::content_blocks_name_and_system(BLOCKS_AND_SYSTEM, 45)] -#[case::openai_tools_named_choice(TOOLS_OPENAI, 123)] -#[case::anthropic_tools_system_discount_choice_none(TOOLS_ANTHROPIC_SYSTEM, 53)] -#[case::completions_prompt(COMPLETIONS_PROMPT, 7)] -#[case::completions_prompt_list(COMPLETIONS_PROMPT_LIST, 4)] -#[case::responses_input_items(RESPONSES_INPUT, 62)] -#[case::embeddings_token_ids(EMBEDDINGS_TOKEN_IDS, 5)] -#[case::rerank_query_and_documents(RERANK, 41)] -fn count_request_matches_python_token_counter(#[case] body: &str, #[case] expected: usize) { - let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses"); - let count = counter().count_request(&request).expect("fixture counts"); - assert_eq!( - count, - InputTokenCount { - model: Some("claude-sonnet-4-5".to_string()), - input_tokens: expected, - } - ); -} + fn assert_count_request_matches_python_token_counter( + load: JsonLoader, + body: &str, + expected: usize, + ) { + let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses"); + let count = counter(load) + .count_request(&request) + .expect("fixture counts"); + assert_eq!( + count, + InputTokenCount { + model: Some("claude-sonnet-4-5".to_string()), + input_tokens: expected, + } + ); + } -#[rstest] -#[case::null_messages_win_over_prompt(r#"{"model":"m","messages":null,"prompt":"ignored"}"#, 3)] -#[case::model_from_route(r#"{"prompt":"hi"}"#, 1)] -#[case::bools_and_ints_use_python_str(r#"{"model":"m","prompt":[true,false,42]}"#, 3)] -#[case::null_prompt_counts_zero(r#"{"model":"m","prompt":null}"#, 0)] -fn key_presence_follows_python(#[case] body: &str, #[case] expected: usize) { - let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses"); - let count = counter().count_request(&request).expect("fixture counts"); - assert_eq!(count.input_tokens, expected); + fn assert_key_presence_follows_python(load: JsonLoader, body: &str, expected: usize) { + let request = CountableRequest::parse(body.as_bytes()).expect("fixture parses"); + let count = counter(load) + .count_request(&request) + .expect("fixture counts"); + assert_eq!(count.input_tokens, expected); + } + + fn assert_shapes_outside_the_mirror_are_declined_at_count(load: JsonLoader, body: &[u8]) { + let request = CountableRequest::parse(body).expect("shape parses"); + assert!(matches!( + counter(load).count_request(&request), + Err(Error::MissingInput | Error::FloatText | Error::ContentBlock | Error::ArrayItems) + )); + } + + fn assert_tool_choice_and_system_discount_change_the_count(load: JsonLoader) { + let counter = counter(load); + let count = |body: &str| { + counter + .count_request(&CountableRequest::parse(body.as_bytes()).expect("parses")) + .expect("counts") + .input_tokens + }; + let base = count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#); + assert_eq!( + count( + r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"none"}"# + ), + base + 1 + ); + assert_eq!( + count( + r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"auto"}"# + ), + base + ); + let with_tools = count( + r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[{"name":"f"}]}"#, + ); + let with_tools_and_system = count( + r#"{"model":"m","messages":[{"role":"system","content":"hi"}],"tools":[{"name":"f"}]}"#, + ); + assert_eq!(with_tools - with_tools_and_system, 4); + assert_eq!( + count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[]}"#), + base + ); + } + + fn assert_loading_a_bad_tokenizer_is_a_load_error(load: JsonLoader) { + assert!(matches!(load("{}"), Err(Error::Load(_)))); + } + + macro_rules! json_backend_tests { + ($loader:path) => { + #[rstest] + #[case::text_only(super::SIMPLE, 14)] + #[case::content_blocks_name_and_system(super::BLOCKS_AND_SYSTEM, 45)] + #[case::openai_tools_named_choice(super::TOOLS_OPENAI, 123)] + #[case::anthropic_tools_system_discount_choice_none(super::TOOLS_ANTHROPIC_SYSTEM, 53)] + #[case::completions_prompt(super::COMPLETIONS_PROMPT, 7)] + #[case::completions_prompt_list(super::COMPLETIONS_PROMPT_LIST, 4)] + #[case::responses_input_items(super::RESPONSES_INPUT, 62)] + #[case::embeddings_token_ids(super::EMBEDDINGS_TOKEN_IDS, 5)] + #[case::rerank_query_and_documents(super::RERANK, 41)] + fn count_request_matches_python_token_counter( + #[case] body: &str, + #[case] expected: usize, + ) { + super::assert_count_request_matches_python_token_counter($loader, body, expected); + } + + #[rstest] + #[case::null_messages_win_over_prompt( + r#"{"model":"m","messages":null,"prompt":"ignored"}"#, + 3 + )] + #[case::model_from_route(r#"{"prompt":"hi"}"#, 1)] + #[case::bools_and_ints_use_python_str(r#"{"model":"m","prompt":[true,false,42]}"#, 3)] + #[case::null_prompt_counts_zero(r#"{"model":"m","prompt":null}"#, 0)] + fn key_presence_follows_python(#[case] body: &str, #[case] expected: usize) { + super::assert_key_presence_follows_python($loader, body, expected); + } + + #[rstest] + #[case::no_countable_input(br#"{"model":"m","instructions":"hi"}"# as &[u8])] + #[case::float_prompt(br#"{"model":"m","prompt":1.5}"#)] + #[case::float_inside_document(br#"{"model":"m","documents":[{"score":0.5}]}"#)] + #[case::image_block( + br#"{"model":"m","messages":[{"role":"user","content":[{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}}]}]}"# + )] + #[case::tool_result_block( + br#"{"model":"m","messages":[{"role":"user","content":[{"type":"tool_result","tool_use_id":"1","content":"ok"}]}]}"# + )] + #[case::array_without_items( + br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"array"}}}}]}"# + )] + fn shapes_outside_the_mirror_are_declined_at_count(#[case] body: &[u8]) { + super::assert_shapes_outside_the_mirror_are_declined_at_count($loader, body); + } + + #[test] + fn tool_choice_and_system_discount_change_the_count() { + super::assert_tool_choice_and_system_discount_change_the_count($loader); + } + + #[test] + fn encoding_errors_preserve_the_backend_source() { + use std::error::Error as _; + + let tokenizer = tokenizers::Tokenizer::new( + tokenizers::models::wordpiece::WordPiece::default(), + ); + let expected = tokenizer.encode_fast("hello", true).unwrap_err(); + let counter = $loader(&tokenizer.to_string(false).unwrap()).unwrap(); + let request = CountableRequest::parse(br#"{"prompt":"hello"}"#).unwrap(); + let error = counter.count_request(&request).unwrap_err(); + assert!(matches!(error, Error::Encode(_))); + assert_eq!(error.source().unwrap().to_string(), expected.to_string()); + } + + #[test] + fn loading_a_bad_tokenizer_is_a_load_error() { + super::assert_loading_a_bad_tokenizer_is_a_load_error($loader); + } + }; + } + + #[cfg(feature = "fast")] + mod fast_json { + use super::*; + + json_backend_tests!(TokenCounter::from_json_fast); + } + + #[cfg(feature = "huggingface")] + mod huggingface_json { + use super::*; + + json_backend_tests!(TokenCounter::from_json); + } } #[rstest] @@ -119,203 +254,204 @@ fn shapes_outside_the_mirror_are_declined_at_parse(#[case] body: &[u8]) { )); } -#[rstest] -#[case::no_countable_input(br#"{"model":"m","instructions":"hi"}"# as &[u8])] -#[case::float_prompt(br#"{"model":"m","prompt":1.5}"#)] -#[case::float_inside_document(br#"{"model":"m","documents":[{"score":0.5}]}"#)] -#[case::image_block( - br#"{"model":"m","messages":[{"role":"user","content":[{"type":"image","source":{"type":"base64","media_type":"image/png","data":"AA=="}}]}]}"# -)] -#[case::tool_result_block( - br#"{"model":"m","messages":[{"role":"user","content":[{"type":"tool_result","tool_use_id":"1","content":"ok"}]}]}"# -)] -#[case::array_without_items( - br#"{"model":"m","messages":[],"tools":[{"name":"f","input_schema":{"type":"object","properties":{"x":{"type":"array"}}}}]}"# -)] -fn shapes_outside_the_mirror_are_declined_at_count(#[case] body: &[u8]) { - let request = CountableRequest::parse(body).expect("shape parses"); +#[cfg(any(feature = "fast", feature = "tiktoken"))] +mod tiktoken { + use super::*; + use serde::Deserialize; + + /// A tiktoken encoding: its fixture directory, the vendored rank file Python + /// loads, the constructor, and the model `generate.py` counted the requests for. + #[derive(Clone, Copy)] + struct TiktokenEncoding { + fixtures: &'static str, + source: TokenizerSource, + load: fn(&str) -> Result, + model: &'static str, + } + + #[cfg(feature = "fast")] + const CL100K: TiktokenEncoding = TiktokenEncoding { + fixtures: "cl100k", + source: TokenizerSource::RankFile("9b5ad71b2ce5302211f9c61530b329a4922fc6a4"), + load: TokenCounter::from_cl100k_ranks, + model: "gpt-4", + }; + + #[cfg(feature = "fast")] + const O200K: TiktokenEncoding = TiktokenEncoding { + fixtures: "o200k", + source: TokenizerSource::RankFile("fb374d419588a4632f3f557e76b4b70aebbca790"), + load: TokenCounter::from_o200k_ranks, + model: "gpt-4o", + }; + + #[cfg(feature = "tiktoken")] + const TIKTOKEN_CL100K: TiktokenEncoding = TiktokenEncoding { + fixtures: "cl100k", + source: TokenizerSource::Name("cl100k_base"), + load: TokenCounter::from_tiktoken, + model: "gpt-4", + }; + + #[cfg(feature = "tiktoken")] + const TIKTOKEN_O200K: TiktokenEncoding = TiktokenEncoding { + fixtures: "o200k", + source: TokenizerSource::Name("o200k_base"), + load: TokenCounter::from_tiktoken, + model: "gpt-4o", + }; + + #[derive(Clone, Copy)] + enum TokenizerSource { + #[cfg(feature = "fast")] + RankFile(&'static str), + #[cfg(feature = "tiktoken")] + Name(&'static str), + } + + fn tiktoken_counter(encoding: TiktokenEncoding) -> TokenCounter { + match encoding.source { + #[cfg(feature = "tiktoken")] + TokenizerSource::Name(name) => (encoding.load)(name).expect("encoding loads"), + #[cfg(feature = "fast")] + TokenizerSource::RankFile(file) => { + let path = format!( + "{}/../../../litellm/litellm_core_utils/tokenizers/{file}", + env!("CARGO_MANIFEST_DIR"), + ); + let ranks = std::fs::read_to_string(path).expect("rank file is in the repo"); + (encoding.load)(&ranks).expect("ranks load") + } + } + } + + fn tiktoken_fixture(encoding: TiktokenEncoding, name: &str) -> String { + let path = format!( + "{}/../token-counter-fast/tests/fixtures/{}/{name}", + env!("CARGO_MANIFEST_DIR"), + encoding.fixtures + ); + std::fs::read_to_string(&path) + .expect("fixture generated by token-counter-fast/tests/fixtures/generate.py") + } + + #[derive(Deserialize)] + struct TextFixture { + text: String, + tokens: usize, + } + + #[derive(Deserialize)] + struct RequestFixture { + body: String, + input_tokens: usize, + } + + /// Reference counts come from `tiktoken.get_encoding(name)`; see + /// `token-counter-fast/tests/fixtures/generate.py`. + #[rstest] + #[cfg_attr(feature = "fast", case::fast_cl100k(CL100K))] + #[cfg_attr(feature = "fast", case::fast_o200k(O200K))] + #[cfg_attr(feature = "tiktoken", case::tiktoken_cl100k(TIKTOKEN_CL100K))] + #[cfg_attr(feature = "tiktoken", case::tiktoken_o200k(TIKTOKEN_O200K))] + fn tiktoken_text_counts_match_tiktoken(#[case] encoding: TiktokenEncoding) { + let counter = tiktoken_counter(encoding); + let fixtures: Vec = tiktoken_fixture(encoding, "texts.jsonl") + .lines() + .map(|line| serde_json::from_str(line).expect("fixture line is json")) + .collect(); + assert!(fixtures.len() > 3000); + let mismatches: Vec<_> = fixtures + .iter() + .filter_map(|fixture| { + let count = counter.count_text(&fixture.text).expect("text counts"); + (count != fixture.tokens).then(|| (fixture.text.clone(), fixture.tokens, count)) + }) + .collect(); + assert!( + mismatches.is_empty(), + "(text, tiktoken, rust): {mismatches:?}" + ); + } + + /// Reference counts come from the proxy's admission counter + /// (`_count_input_tokens(body, model)`), so this pins the shared message, + /// tool and reply-priming accounting on the tiktoken paths as well. + #[rstest] + #[cfg_attr(feature = "fast", case::fast_cl100k(CL100K))] + #[cfg_attr(feature = "fast", case::fast_o200k(O200K))] + #[cfg_attr(feature = "tiktoken", case::tiktoken_cl100k(TIKTOKEN_CL100K))] + #[cfg_attr(feature = "tiktoken", case::tiktoken_o200k(TIKTOKEN_O200K))] + fn tiktoken_request_counts_match_python_admission_counter(#[case] encoding: TiktokenEncoding) { + let counter = tiktoken_counter(encoding); + let fixtures: Vec = tiktoken_fixture(encoding, "requests.jsonl") + .lines() + .map(|line| serde_json::from_str(line).expect("fixture line is json")) + .collect(); + let counts: Vec = fixtures + .iter() + .map(|fixture| { + let request = + CountableRequest::parse(fixture.body.as_bytes()).expect("fixture parses"); + let count = counter.count_request(&request).expect("fixture counts"); + assert_eq!(count.model.as_deref(), Some(encoding.model)); + assert_eq!(count.input_tokens, fixture.input_tokens, "{}", fixture.body); + count.input_tokens + }) + .collect(); + assert!(counts.last().is_some_and(|tokens| *tokens >= 50_000)); + } + + #[rstest] + #[cfg_attr(feature = "fast", case::fast_cl100k(CL100K))] + #[cfg_attr(feature = "fast", case::fast_o200k(O200K))] + #[cfg_attr(feature = "tiktoken", case::tiktoken_cl100k(TIKTOKEN_CL100K))] + #[cfg_attr(feature = "tiktoken", case::tiktoken_o200k(TIKTOKEN_O200K))] + fn tiktoken_shares_the_message_accounting_with_the_anthropic_path( + #[case] encoding: TiktokenEncoding, + ) { + let counter = tiktoken_counter(encoding); + let count = |body: &str| { + counter + .count_request(&CountableRequest::parse(body.as_bytes()).expect("parses")) + .expect("counts") + .input_tokens + }; + let text = |text: &str| counter.count_text(text).expect("counts"); + let base = count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#); + assert_eq!(base, 3 + text("user") + text("hi") + 3); + assert_eq!( + count(r#"{"model":"m","messages":[{"role":"user","name":"al","content":"hi"}]}"#), + base + text("al") + 1 + ); + assert_eq!( + count( + r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"none"}"# + ), + base + 1 + ); + } + + #[cfg(feature = "fast")] + #[rstest] + #[case::empty("")] + #[case::not_base64("!!!! 0")] + #[case::missing_rank("YQ==")] + #[case::rank_not_a_number("YQ== x")] + #[case::single_byte_tokens_missing("YWI= 0")] + fn loading_a_bad_rank_file_is_a_load_error( + #[case] rank_file: &str, + #[values(CL100K, O200K)] encoding: TiktokenEncoding, + ) { + assert!(matches!((encoding.load)(rank_file), Err(Error::Ranks(_)))); + } +} + +#[cfg(feature = "tiktoken")] +#[test] +fn unsupported_encoding_reaches_the_counter_caller() { assert!(matches!( - counter().count_request(&request), - Err(Error::MissingInput | Error::FloatText | Error::ContentBlock | Error::ArrayItems) + TokenCounter::from_tiktoken("unknown-encoding"), + Err(Error::UnsupportedTokenizer(name)) if name == "unknown-encoding" )); } - -#[test] -fn tool_choice_and_system_discount_change_the_count() { - let counter = counter(); - let count = |body: &str| { - counter - .count_request(&CountableRequest::parse(body.as_bytes()).expect("parses")) - .expect("counts") - .input_tokens - }; - let base = count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#); - assert_eq!( - count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"none"}"#), - base + 1 - ); - assert_eq!( - count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"auto"}"#), - base - ); - let with_tools = count( - r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[{"name":"f"}]}"#, - ); - let with_tools_and_system = count( - r#"{"model":"m","messages":[{"role":"system","content":"hi"}],"tools":[{"name":"f"}]}"#, - ); - assert_eq!(with_tools - with_tools_and_system, 4); - assert_eq!( - count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tools":[]}"#), - base - ); -} - -#[test] -fn loading_a_bad_tokenizer_is_a_load_error() { - assert!(matches!(TokenCounter::from_json("{}"), Err(Error::Load(_)))); -} - -/// A tiktoken encoding: its fixture directory, the vendored rank file Python -/// loads, the constructor, and the model `generate.py` counted the requests for. -#[derive(Clone, Copy)] -struct TiktokenEncoding { - fixtures: &'static str, - rank_file: &'static str, - load: fn(&str) -> Result, - model: &'static str, -} - -const CL100K: TiktokenEncoding = TiktokenEncoding { - fixtures: "cl100k", - rank_file: "9b5ad71b2ce5302211f9c61530b329a4922fc6a4", - load: TokenCounter::from_cl100k_ranks, - model: "gpt-4", -}; - -const O200K: TiktokenEncoding = TiktokenEncoding { - fixtures: "o200k", - rank_file: "fb374d419588a4632f3f557e76b4b70aebbca790", - load: TokenCounter::from_o200k_ranks, - model: "gpt-4o", -}; - -fn tiktoken_counter(encoding: TiktokenEncoding) -> TokenCounter { - let path = format!( - "{}/../../../litellm/litellm_core_utils/tokenizers/{}", - env!("CARGO_MANIFEST_DIR"), - encoding.rank_file - ); - let ranks = std::fs::read_to_string(&path).expect("rank file is in the repo"); - (encoding.load)(&ranks).expect("ranks load") -} - -fn tiktoken_fixture(encoding: TiktokenEncoding, name: &str) -> String { - let path = format!( - "{}/tests/fixtures/{}/{name}", - env!("CARGO_MANIFEST_DIR"), - encoding.fixtures - ); - std::fs::read_to_string(&path).expect("fixture generated by tests/fixtures/generate.py") -} - -#[derive(Deserialize)] -struct TextFixture { - text: String, - tokens: usize, -} - -#[derive(Deserialize)] -struct RequestFixture { - body: String, - input_tokens: usize, -} - -/// Reference counts come from `tiktoken.get_encoding(name)`; see -/// `tests/fixtures/generate.py`. -#[rstest] -#[case::cl100k(CL100K)] -#[case::o200k(O200K)] -fn tiktoken_text_counts_match_tiktoken(#[case] encoding: TiktokenEncoding) { - let counter = tiktoken_counter(encoding); - let fixtures: Vec = tiktoken_fixture(encoding, "texts.jsonl") - .lines() - .map(|line| serde_json::from_str(line).expect("fixture line is json")) - .collect(); - assert!(fixtures.len() > 3000); - let mismatches: Vec<_> = fixtures - .iter() - .filter_map(|fixture| { - let count = counter.count_text(&fixture.text).expect("text counts"); - (count != fixture.tokens).then(|| (fixture.text.clone(), fixture.tokens, count)) - }) - .collect(); - assert!( - mismatches.is_empty(), - "(text, tiktoken, rust): {mismatches:?}" - ); -} - -/// Reference counts come from the proxy's admission counter -/// (`_count_input_tokens(body, model)`), so this pins the shared message, -/// tool and reply-priming accounting on the tiktoken paths as well. -#[rstest] -#[case::cl100k(CL100K)] -#[case::o200k(O200K)] -fn tiktoken_request_counts_match_python_admission_counter(#[case] encoding: TiktokenEncoding) { - let counter = tiktoken_counter(encoding); - let fixtures: Vec = tiktoken_fixture(encoding, "requests.jsonl") - .lines() - .map(|line| serde_json::from_str(line).expect("fixture line is json")) - .collect(); - let counts: Vec = fixtures - .iter() - .map(|fixture| { - let request = CountableRequest::parse(fixture.body.as_bytes()).expect("fixture parses"); - let count = counter.count_request(&request).expect("fixture counts"); - assert_eq!(count.model.as_deref(), Some(encoding.model)); - assert_eq!(count.input_tokens, fixture.input_tokens, "{}", fixture.body); - count.input_tokens - }) - .collect(); - assert!(counts.last().is_some_and(|tokens| *tokens >= 50_000)); -} - -#[rstest] -#[case::cl100k(CL100K)] -#[case::o200k(O200K)] -fn tiktoken_shares_the_message_accounting_with_the_anthropic_path( - #[case] encoding: TiktokenEncoding, -) { - let counter = tiktoken_counter(encoding); - let count = |body: &str| { - counter - .count_request(&CountableRequest::parse(body.as_bytes()).expect("parses")) - .expect("counts") - .input_tokens - }; - let text = |text: &str| counter.count_text(text).expect("counts"); - let base = count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}]}"#); - assert_eq!(base, 3 + text("user") + text("hi") + 3); - assert_eq!( - count(r#"{"model":"m","messages":[{"role":"user","name":"al","content":"hi"}]}"#), - base + text("al") + 1 - ); - assert_eq!( - count(r#"{"model":"m","messages":[{"role":"user","content":"hi"}],"tool_choice":"none"}"#), - base + 1 - ); -} - -#[rstest] -#[case::empty("")] -#[case::not_base64("!!!! 0")] -#[case::missing_rank("YQ==")] -#[case::rank_not_a_number("YQ== x")] -#[case::single_byte_tokens_missing("YWI= 0")] -fn loading_a_bad_rank_file_is_a_load_error( - #[case] rank_file: &str, - #[values(CL100K, O200K)] encoding: TiktokenEncoding, -) { - assert!(matches!((encoding.load)(rank_file), Err(Error::Ranks(_)))); -} diff --git a/litellm/__init__.py b/litellm/__init__.py index 738dd0cac76..be8f59d210b 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -400,6 +400,7 @@ default_redis_batch_cache_expiry: Optional[float] = None model_alias_map: Dict[str, str] = {} model_group_settings: Optional["ModelGroupSettings"] = None max_budget: float = 0.0 # set the max budget across all providers +budget_exceeded_status_code: int = 422 # set to 429 to restore the pre-422 budget_exceeded response code budget_duration: Optional[str] = ( None # proxy only - resets budget after fixed duration. You can set duration as seconds ("30s"), minutes ("30m"), hours ("30h"), days ("30d"). ) diff --git a/litellm/_logging.py b/litellm/_logging.py index 5ba0c080364..644a79d8cbd 100644 --- a/litellm/_logging.py +++ b/litellm/_logging.py @@ -5,6 +5,7 @@ import logging import os import re import sys +from collections.abc import Sequence from datetime import datetime from logging import Formatter from typing import Any, Final, TextIO @@ -186,7 +187,8 @@ class SecretRedactionFilter(logging.Filter): record.stack_info = _redact_string(record.stack_info) # rebind-ok: a Filter scrubs records in place # Redact extra fields passed via logger.debug("msg", extra={...}) - for key, value in list(record.__dict__.items()): + record_items: Final[Sequence[tuple[str, object]]] = list(record.__dict__.items()) + for key, value in record_items: if key in _STANDARD_RECORD_ATTRS: continue if isinstance(value, str): @@ -507,7 +509,7 @@ handler.addFilter(_secret_filter) handler.addFilter(_correlation_filter) -def _try_parse_json_message(message: str) -> dict[str, Any] | None: +def _try_parse_json_message(message: str) -> dict[str, object] | None: """ Try to parse a log message as JSON. Returns parsed dict if valid, else None. Handles messages that are entirely valid JSON (e.g. json.dumps output). @@ -585,7 +587,7 @@ class JsonFormatter(Formatter): def format(self, record): message_str: Final = record.getMessage() - json_record: Final[dict[str, Any]] = { + json_record: Final[dict[str, object]] = { "message": message_str, "level": record.levelname, "timestamp": self.formatTime(record), diff --git a/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py index 306a8871b12..da5eb522187 100644 --- a/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py +++ b/litellm/a2a_protocol/providers/bedrock_agentcore/handler.py @@ -98,7 +98,7 @@ class BedrockAgentCoreA2AHandler: request_id=request_id, params=params, litellm_params=litellm_params, - method="message/send", + method="message/stream", stream=True, agent_extra_headers=agent_extra_headers, ) diff --git a/litellm/a2a_protocol/streaming_iterator.py b/litellm/a2a_protocol/streaming_iterator.py index d936caeb75e..8232d7cf2d8 100644 --- a/litellm/a2a_protocol/streaming_iterator.py +++ b/litellm/a2a_protocol/streaming_iterator.py @@ -5,7 +5,7 @@ A2A Streaming Iterator with token tracking and logging support. import asyncio from collections.abc import AsyncIterator from datetime import datetime -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Final import litellm from litellm._logging import verbose_logger @@ -15,7 +15,7 @@ from litellm.litellm_core_utils.asyncify import asyncify from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj if TYPE_CHECKING: - from a2a.types import SendStreamingMessageRequest, SendStreamingMessageResponse + from a2a.compat.v0_3.types import SendStreamingMessageRequest, SendStreamingMessageResponse class A2AStreamingIterator: @@ -39,9 +39,9 @@ class A2AStreamingIterator: self.start_time = datetime.now() # Collect chunks for token counting - self.chunks: list[Any] = [] + self.chunks: list[SendStreamingMessageResponse] = [] self.collected_text_parts: list[str] = [] - self.final_chunk: Any | None = None + self.final_chunk: SendStreamingMessageResponse | None = None def __aiter__(self): return self @@ -69,7 +69,7 @@ class A2AStreamingIterator: await self._handle_stream_complete() raise - def _collect_text_from_chunk(self, chunk: Any) -> None: + def _collect_text_from_chunk(self, chunk: "SendStreamingMessageResponse") -> None: """Extract text from a streaming chunk and add to collected parts.""" try: chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {} @@ -79,7 +79,7 @@ class A2AStreamingIterator: except Exception: verbose_logger.debug("Failed to extract text from A2A streaming chunk") - def _is_completed_chunk(self, chunk: Any) -> bool: + def _is_completed_chunk(self, chunk: "SendStreamingMessageResponse") -> bool: """Check if chunk indicates stream completion.""" try: chunk_dict: Final = chunk.model_dump(mode="json", exclude_none=True) if hasattr(chunk, "model_dump") else {} diff --git a/litellm/exceptions.py b/litellm/exceptions.py index 14cc16452f0..c8de2ab12ed 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -16,6 +16,7 @@ from typing import Any, Final import httpx import openai +import litellm from litellm.types.utils import LiteLLMCommonStrings from litellm.types.vector_stores import VectorStoreSearchFailure @@ -1002,7 +1003,7 @@ class BudgetExceededError(Exception): ): self.current_cost = current_cost self.max_budget = max_budget - self.status_code = 429 + self.status_code = litellm.budget_exceeded_status_code self.llm_provider = llm_provider or "" self.entity_type = entity_type self.entity_id = entity_id diff --git a/litellm/experimental_mcp_client/Readme.md b/litellm/experimental_mcp_client/Readme.md index 14decce0256..3385a37cf69 100644 --- a/litellm/experimental_mcp_client/Readme.md +++ b/litellm/experimental_mcp_client/Readme.md @@ -2,6 +2,17 @@ LiteLLM MCP Client allows you to use MCP tools with LiteLLM +Install the optional dependencies with `pip install 'litellm[mcp]'`, then use the existing public imports: + +```python +from litellm.experimental_mcp_client import call_openai_tool, load_mcp_tools +from litellm.experimental_mcp_client.client import MCPClient + +client = MCPClient(server_url="https://mcp.example.com/mcp") +``` + +Core `import litellm` works without the MCP extra. Importing the experimental MCP client without its MCP or HTTPX2 dependency raises an error with this installation command + ## MCP Python SDK compatibility The `mcp` and `proxy` extras require MCP Python SDK 2.2 or newer within the 2.x release line. Installing core LiteLLM without these extras does not require MCP @@ -16,6 +27,12 @@ The shared unit-test workflow runs the MCP integration suite once, with SDK2 in See the official [SDK migration guide](https://py.sdk.modelcontextprotocol.io/migration/) for Python API changes +## Custom HTTP clients and authentication + +MCP HTTP and SSE transports now use `httpx2`. Custom authentication passed through `aws_auth` or `resolved_auth` must implement `httpx2.Auth`. Integrations that override the client's HTTP client factory or customize its event hooks must use `httpx2.AsyncClient`, request, response, timeout and transport types + +HTTPX1 clients, auth objects and hooks are not adapted by a compatibility shim. Migrate those integrations to HTTPX2 before upgrading. Ordinary `MCPClient` construction and LiteLLM's existing helper imports remain supported; this does not restore SDK1 Python imports or camelCase SDK model attributes in the shared Python environment + ## HTTP redirects For streamable HTTP POST requests, the MCP SDK follows method-preserving redirects such as HTTP 307/308 within the configured endpoint's origin. Redirects to another path on the same scheme, host and port work. The SDK also permits an HTTP-to-HTTPS upgrade on the same host using the default ports diff --git a/litellm/experimental_mcp_client/__init__.py b/litellm/experimental_mcp_client/__init__.py index 5399968ff74..7a3917a50ea 100644 --- a/litellm/experimental_mcp_client/__init__.py +++ b/litellm/experimental_mcp_client/__init__.py @@ -1,3 +1,8 @@ -from .tools import call_openai_tool, load_mcp_tools +try: + from .tools import call_openai_tool, load_mcp_tools +except ModuleNotFoundError as exc: + if exc.name not in ("mcp", "httpx2"): + raise + raise ImportError("MCP client dependencies are missing. Install them with: pip install 'litellm[mcp]'") from exc __all__ = ["call_openai_tool", "load_mcp_tools"] diff --git a/litellm/integrations/focus/focus_logger.py b/litellm/integrations/focus/focus_logger.py index c9b47835948..dce51b8190b 100644 --- a/litellm/integrations/focus/focus_logger.py +++ b/litellm/integrations/focus/focus_logger.py @@ -15,6 +15,8 @@ from .destinations import FocusTimeWindow if TYPE_CHECKING: from apscheduler.schedulers.asyncio import AsyncIOScheduler + from litellm.proxy.db.db_transaction_queue.pod_lock_manager import PodLockManager + from .export_engine import FocusExportEngine else: AsyncIOScheduler = Any @@ -111,7 +113,7 @@ class FocusLogger(CustomLogger): """Entry point for scheduler jobs to run export cycle with locking.""" from litellm.proxy.proxy_server import proxy_logging_obj - pod_lock_manager = None + pod_lock_manager: PodLockManager | None = None if proxy_logging_obj is not None: writer: Final = getattr(proxy_logging_obj, "db_spend_update_writer", None) if writer is not None: diff --git a/litellm/integrations/generic_prompt_management/generic_prompt_manager.py b/litellm/integrations/generic_prompt_management/generic_prompt_manager.py index 77d315d0cee..de01b2bb02c 100644 --- a/litellm/integrations/generic_prompt_management/generic_prompt_manager.py +++ b/litellm/integrations/generic_prompt_management/generic_prompt_manager.py @@ -58,7 +58,7 @@ class GenericPromptManager(CustomPromptManagement): api_key: str | None = None, timeout: int = 30, prompt_id: str | None = None, - additional_provider_specific_query_params: dict[str, Any] | None = None, + additional_provider_specific_query_params: Mapping[str, object] | None = None, **kwargs, ): """ diff --git a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py index 3c189b4d53e..7e3c4cc3ce8 100644 --- a/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py +++ b/litellm/integrations/mavvrik_focus/mavvrik_focus_logger.py @@ -21,7 +21,7 @@ from __future__ import annotations import os from datetime import datetime, timedelta, timezone -from typing import TYPE_CHECKING, Any, Final +from typing import TYPE_CHECKING, Any, Final, Protocol import litellm from litellm._logging import verbose_proxy_logger @@ -35,6 +35,17 @@ else: AsyncIOScheduler = Any +class _PodLockManager(Protocol): + """The subset of PodLockManager this logger drives to serialize the export across pods.""" + + @property + def redis_cache(self) -> object: ... + + async def acquire_lock(self, cronjob_id: str) -> bool | None: ... + + async def release_lock(self, cronjob_id: str) -> None: ... + + def _parse_metrics_marker( marker: object | None, ) -> datetime | None: @@ -226,9 +237,9 @@ class MavvrikFocusLogger(FocusLogger): """Scheduler entry point — uses Mavvrik-specific pod-lock key.""" from litellm.proxy.proxy_server import proxy_logging_obj # noqa: PLC0415 - pod_lock_manager = None + pod_lock_manager: _PodLockManager | None = None if proxy_logging_obj is not None: - writer: Final = getattr(proxy_logging_obj, "db_spend_update_writer", None) + writer: Final[object] = getattr(proxy_logging_obj, "db_spend_update_writer", None) if writer is not None: pod_lock_manager = getattr(writer, "pod_lock_manager", None) diff --git a/litellm/integrations/otel/presets/agentops.py b/litellm/integrations/otel/presets/agentops.py index 965213f2ee4..58123656caa 100644 --- a/litellm/integrations/otel/presets/agentops.py +++ b/litellm/integrations/otel/presets/agentops.py @@ -9,9 +9,12 @@ this preset registers a custom exporter (``kind="agentops"``) that mints the JWT worker thread, off any event loop — and caches it for the process lifetime. """ +from collections.abc import Sequence from typing import Any, Final import httpx +from opentelemetry.sdk.trace import ReadableSpan +from opentelemetry.sdk.trace.export import SpanExporter, SpanExportResult from pydantic import Field from pydantic_settings import BaseSettings, SettingsConfigDict @@ -71,7 +74,7 @@ def agentops_preset( ) -def _build_agentops_exporter(spec: ExporterSpec) -> Any: +def _build_agentops_exporter(spec: ExporterSpec) -> SpanExporter: """Factory for the ``agentops`` exporter kind: a lazy-auth OTLP/HTTP exporter.""" from opentelemetry.exporter.otlp.proto.http.trace_exporter import ( OTLPSpanExporter, @@ -106,7 +109,7 @@ def _build_agentops_exporter(spec: ExporterSpec) -> Any: except Exception as e: verbose_logger.debug("AgentOps JWT fetch failed: %s", e) - def export(self, spans: Any) -> Any: + def export(self, spans: Sequence[ReadableSpan]) -> SpanExportResult: self._ensure_authenticated() return super().export(spans) diff --git a/litellm/integrations/otel/runtime.py b/litellm/integrations/otel/runtime.py index c6eaecd108b..13903597e1a 100644 --- a/litellm/integrations/otel/runtime.py +++ b/litellm/integrations/otel/runtime.py @@ -8,13 +8,16 @@ identity unconditionally. """ from collections.abc import Callable, Iterator -from contextlib import contextmanager +from contextlib import AbstractContextManager, contextmanager from functools import cache -from typing import Any, Final +from typing import TYPE_CHECKING, Final + +if TYPE_CHECKING: + from opentelemetry.trace import Span @cache -def _otel_runtime() -> "tuple[Callable[[str], Any], Callable[..., None]] | None": +def _otel_runtime() -> "tuple[Callable[[str], AbstractContextManager[Span | None]], Callable[..., None]] | None": """Resolve the SDK-backed hooks once and cache the outcome, absence included. CPython never caches a failed import, so without this memoization every call @@ -29,7 +32,7 @@ def _otel_runtime() -> "tuple[Callable[[str], Any], Callable[..., None]] | None" @contextmanager -def phase_span(name: str) -> "Iterator[Any]": +def phase_span(name: str) -> "Iterator[Span | None]": """Run a request phase inside a live active span so its DB/service calls nest. Yields ``None`` (a plain no-op) when the OTel SDK is unavailable or V2 is not @@ -43,7 +46,7 @@ def phase_span(name: str) -> "Iterator[Any]": yield span -def seed_request_identity(user_api_key_dict: Any, model: Any = None) -> None: +def seed_request_identity(user_api_key_dict: object, model: object = None) -> None: """Seed request-identity Baggage at the auth boundary (no-op without V2).""" runtime: Final = _otel_runtime() if runtime is None: diff --git a/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py b/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py index c1ccf09d5d6..ba7d54fafea 100644 --- a/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py +++ b/litellm/integrations/prometheus_helpers/bounded_prometheus_series_tracker.py @@ -3,7 +3,13 @@ from __future__ import annotations import time from collections import OrderedDict from threading import RLock -from typing import Any, Final +from typing import Final, Protocol + + +class _RemovableMetric(Protocol): + """The one prometheus-client metric method this tracker calls.""" + + def remove(self, *labelvalues: object) -> None: ... class BoundedPrometheusSeriesTracker: @@ -21,7 +27,7 @@ class BoundedPrometheusSeriesTracker: def track_series( self, - metric: Any, + metric: _RemovableMetric, metric_name: str, label_values: tuple[str | None, ...], max_series: int | None, @@ -60,7 +66,7 @@ class BoundedPrometheusSeriesTracker: break del series[tracked_label_values] - def remove_series(self, metric: object, label_values: tuple[str | None, ...]) -> bool: + def remove_series(self, metric: _RemovableMetric, label_values: tuple[str | None, ...]) -> bool: """Drop one child series, True when it is gone (removed or never existed).""" return self._remove_metric_child(metric, label_values) @@ -82,7 +88,7 @@ class BoundedPrometheusSeriesTracker: def _remove_metric_series( self, - metric: Any, + metric: _RemovableMetric, series: OrderedDict[tuple[str | None, ...], float], label_values: tuple[str | None, ...], ) -> None: @@ -90,7 +96,7 @@ class BoundedPrometheusSeriesTracker: series.pop(label_values, None) @staticmethod - def _remove_metric_child(metric: Any, label_values: tuple[str | None, ...]) -> bool: + def _remove_metric_child(metric: _RemovableMetric, label_values: tuple[str | None, ...]) -> bool: """ Remove the Prometheus child for ``label_values`` and report whether the tracker should commit the matching state change. diff --git a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py index b2243060c6c..74fb8a8d6a3 100644 --- a/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py +++ b/litellm/integrations/vector_store_integrations/vector_store_pre_call_hook.py @@ -406,7 +406,7 @@ class VectorStorePreCallHook(CustomLogger): request_data: dict, response_chunk: Any, call_type: CallTypes | None, - ) -> Any | None: + ) -> object | None: """ Add search results to the final streaming chunk. diff --git a/litellm/integrations/weights_biases.py b/litellm/integrations/weights_biases.py index d1a8ec098cf..9bc070a1f9a 100644 --- a/litellm/integrations/weights_biases.py +++ b/litellm/integrations/weights_biases.py @@ -4,6 +4,7 @@ imported_openAIResponse = True try: import io import logging + from collections.abc import Mapping from typing import Any, Literal, Protocol, TypeVar from wandb.sdk.data_types import trace_tree @@ -43,7 +44,7 @@ try: @staticmethod def results_to_trace_tree( - request: dict[str, Any], + request: Mapping[str, object], response: OpenAIResponse, results: list[trace_tree.Result], time_elapsed: float, @@ -73,7 +74,7 @@ try: def _resolve_edit( self, - request: dict[str, Any], + request: Mapping[str, object], response: OpenAIResponse, time_elapsed: float, ) -> trace_tree.WBTraceTree: @@ -91,7 +92,7 @@ try: def _resolve_completion( self, - request: dict[str, Any], + request: Mapping[str, object], response: OpenAIResponse, time_elapsed: float, ) -> trace_tree.WBTraceTree: @@ -134,7 +135,7 @@ try: def _request_response_result_to_trace( self, - request: dict[str, Any], + request: Mapping[str, object], response: OpenAIResponse, request_str: str, choices: list[str], diff --git a/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py b/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py index f11f6d46fb2..a02c40b7611 100644 --- a/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py +++ b/litellm/litellm_core_utils/llm_cost_calc/usage_object_transformation.py @@ -50,7 +50,7 @@ _INTERACTIONS_MODALITY_FIELDS: Final[Mapping[str, str]] = MappingProxyType( ) -def _modality_field(entry: Mapping[str, Any]) -> str | None: +def _modality_field(entry: Mapping[str, object]) -> str | None: return _INTERACTIONS_MODALITY_FIELDS.get(str(entry.get("modality", "")).lower()) @@ -58,7 +58,7 @@ def _token_count(value: object) -> int: return value if isinstance(value, int) else 0 -def _modality_token_sums(entries: Sequence[Mapping[str, Any]]) -> Mapping[str, int]: +def _modality_token_sums(entries: Sequence[Mapping[str, object]]) -> Mapping[str, int]: fields: Final = frozenset(field for entry in entries if (field := _modality_field(entry)) is not None) return MappingProxyType( { @@ -68,7 +68,7 @@ def _modality_token_sums(entries: Sequence[Mapping[str, Any]]) -> Mapping[str, i ) -def _google_search_query_count(usage_object: Mapping[str, Any]) -> int: +def _google_search_query_count(usage_object: Mapping[str, object]) -> int: entries: Final = usage_object.get("grounding_tool_count") if not isinstance(entries, Sequence): return 0 diff --git a/litellm/litellm_core_utils/safe_json_dumps.py b/litellm/litellm_core_utils/safe_json_dumps.py index 4f9ac82d57d..63242a580e7 100644 --- a/litellm/litellm_core_utils/safe_json_dumps.py +++ b/litellm/litellm_core_utils/safe_json_dumps.py @@ -85,7 +85,7 @@ def safe_json_structure( def safe_dumps( - data: Any, + data: object, max_depth: int = DEFAULT_MAX_RECURSE_DEPTH, value_transform: Callable[[str | None, str], str] | None = None, ) -> str: diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py b/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py index 87a29ca50ba..54d10837d74 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/handler.py @@ -35,7 +35,7 @@ if TYPE_CHECKING: from litellm.router import Router # Anthropic-only keys already mapped by the translator; strip on extra_kwargs re-merge. -ANTHROPIC_ONLY_REQUEST_KEYS: Final[frozenset[str]] = frozenset({"output_config"}) +ANTHROPIC_ONLY_REQUEST_KEYS: Final[frozenset[str]] = frozenset({"output_config", "safeguards"}) _AnthropicMessages: TypeAlias = "list[dict[str, object]]" _AnthropicSystem: TypeAlias = "str | list[dict[str, object]] | None" diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py index 5fa686b7560..eed30c2698c 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/transformation.py @@ -79,10 +79,14 @@ class AnthropicMessagesConfig(BaseAnthropicMessagesConfig): "speed", "output_config", "reasoning_effort", + "safeguards", # TODO: Add Anthropic `metadata` support # "metadata", ] + def should_filter_anthropic_beta_headers(self) -> bool: + return self._resolved_provider != "anthropic" + def _remove_scope_from_cache_control(self, anthropic_messages_request: dict) -> None: """ Remove `scope` field from cache_control blocks. diff --git a/litellm/llms/azure/fine_tuning/handler.py b/litellm/llms/azure/fine_tuning/handler.py index ac1e430e063..36c4fae04c7 100644 --- a/litellm/llms/azure/fine_tuning/handler.py +++ b/litellm/llms/azure/fine_tuning/handler.py @@ -1,5 +1,5 @@ from collections.abc import Coroutine -from typing import Any, Final, cast +from typing import Final, cast import httpx from openai import AsyncAzureOpenAI, AsyncOpenAI, AzureOpenAI, OpenAI @@ -19,7 +19,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): """ @staticmethod - def _ensure_training_type(create_fine_tuning_job_data: dict[str, Any]) -> None: + def _ensure_training_type(create_fine_tuning_job_data: dict[str, object]) -> None: """ Azure requires trainingType in extra_body. Default to 1 (supervised) if omitted. """ @@ -66,7 +66,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): max_retries: int | None, organization: str | None, client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, - ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: + ) -> LiteLLMFineTuningJob | Coroutine[object, object, LiteLLMFineTuningJob]: self._ensure_training_type(create_fine_tuning_job_data) openai_client: Final[OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None] = self.get_openai_client( @@ -109,7 +109,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): max_retries: int | None, organization: str | None, client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, - ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: + ) -> LiteLLMFineTuningJob | Coroutine[object, object, LiteLLMFineTuningJob]: openai_client: Final[OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None] = self.get_openai_client( api_key=api_key, api_base=api_base, @@ -149,7 +149,7 @@ class AzureOpenAIFineTuningAPI(OpenAIFineTuningAPI, BaseAzureLLM): max_retries: int | None, organization: str | None, client: OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None = None, - ) -> LiteLLMFineTuningJob | Coroutine[Any, Any, LiteLLMFineTuningJob]: + ) -> LiteLLMFineTuningJob | Coroutine[object, object, LiteLLMFineTuningJob]: openai_client: Final[OpenAI | AsyncOpenAI | AzureOpenAI | AsyncAzureOpenAI | None] = self.get_openai_client( api_key=api_key, api_base=api_base, diff --git a/litellm/llms/codestral/completion/transformation.py b/litellm/llms/codestral/completion/transformation.py index e3c3fd1231c..baa134bb398 100644 --- a/litellm/llms/codestral/completion/transformation.py +++ b/litellm/llms/codestral/completion/transformation.py @@ -29,7 +29,7 @@ class CodestralTextCompletionConfig(OpenAITextCompletionConfig): random_seed: int | None = None, stop: str | None = None, ) -> None: - locals_: Final = locals().copy() + locals_: Final[dict[str, object]] = locals().copy() for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/databricks/common_utils.py b/litellm/llms/databricks/common_utils.py index a4ec2c5378b..f2f9df422e8 100644 --- a/litellm/llms/databricks/common_utils.py +++ b/litellm/llms/databricks/common_utils.py @@ -12,7 +12,7 @@ Authentication priority: import os import re -from typing import Any, Final, Literal +from typing import Final, Literal from urllib.parse import urlsplit, urlunsplit from litellm.llms.base_llm.chat.transformation import BaseLLMException @@ -48,7 +48,7 @@ class DatabricksBase: ] @classmethod - def redact_sensitive_data(cls, data: Any) -> Any: + def redact_sensitive_data(cls, data: object) -> object: """ Redact sensitive information (tokens, secrets) from data before logging. diff --git a/litellm/llms/gemini/realtime/transformation.py b/litellm/llms/gemini/realtime/transformation.py index 79985569c5f..e64cbf88d95 100644 --- a/litellm/llms/gemini/realtime/transformation.py +++ b/litellm/llms/gemini/realtime/transformation.py @@ -453,7 +453,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): return normalized @staticmethod - def _finalize_gemini_live_setup(model: str, setup: dict[str, Any]) -> dict[str, Any]: + def _finalize_gemini_live_setup(model: str, setup: dict[str, object]) -> dict[str, object]: generation_config: Final = setup.get("generationConfig") if isinstance(generation_config, dict): modalities: Final = generation_config.get("responseModalities") @@ -1172,7 +1172,7 @@ class GeminiRealtimeConfig(BaseRealtimeConfig): def map_openai_event( self, key: str, - value: Any, + value: object, current_delta_type: ALL_DELTA_TYPES | None, ) -> OpenAIRealtimeEventTypes | ResponsesAPIStreamEvents: if isinstance(value, dict): diff --git a/litellm/llms/jina_ai/embedding/transformation.py b/litellm/llms/jina_ai/embedding/transformation.py index 26f512979e5..260d9e6e494 100644 --- a/litellm/llms/jina_ai/embedding/transformation.py +++ b/litellm/llms/jina_ai/embedding/transformation.py @@ -31,7 +31,7 @@ class JinaAIEmbeddingConfig(BaseEmbeddingConfig): def __init__( self, ) -> None: - locals_: Final = locals().copy() + locals_: Final[dict[str, object]] = locals().copy() for key, value in locals_.items(): if key != "self" and value is not None: setattr(self.__class__, key, value) diff --git a/litellm/llms/openrouter/embedding/transformation.py b/litellm/llms/openrouter/embedding/transformation.py index 29d0c8c1c56..14b0e462ea7 100644 --- a/litellm/llms/openrouter/embedding/transformation.py +++ b/litellm/llms/openrouter/embedding/transformation.py @@ -170,7 +170,9 @@ class OpenrouterEmbeddingConfig(BaseEmbeddingConfig): optional_params[param] = value return optional_params - def get_error_class(self, error_message: str, status_code: int, headers: Any) -> Any: + def get_error_class( + self, error_message: str, status_code: int, headers: dict[str, str] | httpx.Headers + ) -> OpenRouterException: """ Get the error class for OpenRouter errors. """ diff --git a/litellm/llms/reducto/common.py b/litellm/llms/reducto/common.py index 9b9efd24b72..b194590fdb9 100644 --- a/litellm/llms/reducto/common.py +++ b/litellm/llms/reducto/common.py @@ -3,6 +3,8 @@ import binascii from collections import defaultdict from typing import TYPE_CHECKING, Any, Final, NoReturn +import httpx + from litellm.constants import request_timeout REDUCTO_API_BASE: Final = "https://platform.reducto.ai" @@ -62,7 +64,7 @@ def extract_file_id_or_bytes( return None, raw_bytes, mime -def _extract_file_id_from_upload_response(response: Any) -> str: +def _extract_file_id_from_upload_response(response: httpx.Response) -> str: try: payload: Final = response.json() except ValueError as exc: diff --git a/litellm/llms/vercel_ai_gateway/embedding/transformation.py b/litellm/llms/vercel_ai_gateway/embedding/transformation.py index 3f228c0881d..fc9c6bcc19f 100644 --- a/litellm/llms/vercel_ai_gateway/embedding/transformation.py +++ b/litellm/llms/vercel_ai_gateway/embedding/transformation.py @@ -11,6 +11,7 @@ from typing import TYPE_CHECKING, Any, Final import httpx +from litellm.llms.base_llm.chat.transformation import BaseLLMException from litellm.llms.base_llm.embedding.transformation import BaseEmbeddingConfig from litellm.secret_managers.main import get_secret_str from litellm.types.llms.openai import AllEmbeddingInputValues @@ -160,7 +161,7 @@ class VercelAIGatewayEmbeddingConfig(BaseEmbeddingConfig): optional_params[param] = value return optional_params - def get_error_class(self, error_message: str, status_code: int, headers: Any) -> Any: + def get_error_class(self, error_message: str, status_code: int, headers: Any) -> BaseLLMException: """ Get the error class for Vercel AI Gateway errors. """ diff --git a/litellm/llms/vertex_ai/agent_engine/transformation.py b/litellm/llms/vertex_ai/agent_engine/transformation.py index e430d9e2280..c5ca9f38144 100644 --- a/litellm/llms/vertex_ai/agent_engine/transformation.py +++ b/litellm/llms/vertex_ai/agent_engine/transformation.py @@ -205,7 +205,7 @@ class VertexAgentEngineConfig(BaseConfig, VertexBase): session_id: Final = self._get_session_id(optional_params) # Build the input - input_data: Final[dict[str, Any]] = { + input_data: Final[dict[str, str]] = { "message": prompt, "user_id": user_id, } diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index bf65675f3c5..5b686f0a4b5 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -1327,7 +1327,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 2048, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "anthropic.claude-mythos-preview": { "input_cost_per_token": 0, @@ -1381,7 +1381,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 2048, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-opus-4-7": { "bedrock_converse_supports_strict_tools": false, @@ -1419,7 +1419,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 2048, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-opus-4-7": { "bedrock_converse_supports_strict_tools": false, @@ -1531,7 +1531,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "anthropic.claude-fable-5-1": { "cache_creation_input_token_cost": 1.25e-05, @@ -1570,7 +1570,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-fable-5": { "cache_creation_input_token_cost": 1.25e-05, @@ -1608,7 +1608,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-fable-5-1": { "cache_creation_input_token_cost": 1.25e-05, @@ -1647,7 +1647,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-fable-5": { "cache_creation_input_token_cost": 1.375e-05, @@ -1685,7 +1685,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-fable-5-1": { "cache_creation_input_token_cost": 1.375e-05, @@ -1724,7 +1724,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-fable-5": { "cache_creation_input_token_cost": 1.375e-05, @@ -1837,7 +1837,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-opus-5": { "bedrock_converse_supports_strict_tools": false, @@ -1875,7 +1875,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-opus-5": { "bedrock_converse_supports_strict_tools": false, @@ -1913,7 +1913,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-opus-5": { "bedrock_converse_supports_strict_tools": false, @@ -2063,7 +2063,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, @@ -2102,7 +2102,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, @@ -2141,7 +2141,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, @@ -2329,7 +2329,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-sonnet-5": { "bedrock_converse_supports_strict_tools": false, @@ -2368,7 +2368,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-sonnet-5": { "bedrock_converse_supports_strict_tools": false, @@ -2407,7 +2407,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-sonnet-5": { "bedrock_converse_supports_strict_tools": false, @@ -2556,7 +2556,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-sonnet-4-6": { "supports_adaptive_thinking": true, @@ -2591,7 +2591,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-sonnet-4-6": { "supports_adaptive_thinking": true, @@ -2626,7 +2626,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-sonnet-4-6": { "supports_adaptive_thinking": true, @@ -3740,6 +3740,21 @@ "supports_vision": true, "supports_web_search": true }, + "azure_ai/gpt-image-2": { + "cache_read_input_image_token_cost": 2e-06, + "cache_read_input_token_cost": 1.25e-06, + "input_cost_per_image_token": 8e-06, + "input_cost_per_token": 5e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image_token": 3e-05, + "source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ], + "supports_vision": true + }, "azure_ai/codex-mini": { "cache_read_input_token_cost": 3.75e-07, "deprecation_date": "2026-11-15", @@ -11159,6 +11174,20 @@ ], "deprecation_date": "2026-10-01" }, + "azure_ai/MAI-Image-2.5-Pro": { + "deprecation_date": "2026-10-01", + "input_cost_per_image_token": 8e-06, + "input_cost_per_token": 5e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image": 0.1085, + "output_cost_per_image_token": 0.000106, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-image-2-5-pro-and-mai-voice-2-flash-in-microsoft-foundry/4539446", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ] + }, "azure_ai/MAI-Image-2e": { "deprecation_date": "2026-08-15", "input_cost_per_token": 5e-06, @@ -21888,6 +21917,7 @@ "supports_tool_choice": true }, "deepseek/deepseek-coder": { + "cache_read_input_token_cost": 1.4e-08, "input_cost_per_token": 1.4e-07, "input_cost_per_token_cache_hit": 1.4e-08, "litellm_provider": "deepseek", @@ -21902,6 +21932,7 @@ "supports_tool_choice": true }, "deepseek/deepseek-r1": { + "cache_read_input_token_cost": 1.4e-07, "input_cost_per_token": 5.5e-07, "input_cost_per_token_cache_hit": 1.4e-07, "litellm_provider": "deepseek", @@ -21957,6 +21988,7 @@ "supports_tool_choice": true }, "deepseek/deepseek-v3.2": { + "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", @@ -23796,6 +23828,25 @@ "supports_tool_choice": true, "supports_vision": false }, + "fireworks_ai/deepseek-v4-pro-0813": { + "cache_read_input_token_cost": 4.4e-08, + "cache_read_input_token_cost_priority": 5.5e-08, + "input_cost_per_token": 1.32e-06, + "input_cost_per_token_priority": 1.65e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 3.96e-06, + "output_cost_per_token_priority": 4.95e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/accounts/fireworks/models/firefunction-v2": { "input_cost_per_token": 9e-07, "litellm_provider": "fireworks_ai", @@ -24182,7 +24233,7 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/mixtral-8x22b-instruct-hf": { "input_cost_per_token": 1.2e-06, @@ -24508,7 +24559,7 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": false }, "fireworks_ai/qwen3p7-plus": { "cache_read_input_token_cost": 8e-08, @@ -35244,6 +35295,7 @@ "mode": "chat", "output_cost_per_token": 3e-06, "source": "https://console.groq.com/docs/model/qwen/qwen3.6-27b", + "deprecation_date": "2026-09-14", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": false, @@ -41472,6 +41524,7 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v3.2-exp": { + "cache_read_input_token_cost": 2e-08, "deprecation_date": "2026-09-28", "input_cost_per_token": 2.7e-07, "input_cost_per_token_cache_hit": 2e-08, @@ -41494,6 +41547,7 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-r1": { + "cache_read_input_token_cost": 1.4e-07, "input_cost_per_token": 7e-07, "input_cost_per_token_cache_hit": 1.4e-07, "litellm_provider": "openrouter", @@ -41537,21 +41591,21 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 4.22298e-07, + "input_cost_per_token": 9.42906e-07, "input_cost_per_token_cache_hit": 4.4e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 8.44596e-07, + "output_cost_per_token": 1.885812e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 3.51915e-08, + "cache_read_input_token_cost": 7.85755e-08, "supports_audio_input": false, "supports_pdf_input": false, "supports_vision": false, @@ -41579,22 +41633,22 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "input_cost_per_token": 5.3328e-07, - "input_cost_per_token_cache_hit": 4.4e-08, + "input_cost_per_token": 5.7684e-07, + "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.59984e-06, + "output_cost_per_token": 1.73052e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 1.6968e-08, - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"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"}],"input_cost_per_token":5.3328e-7,"output_cost_per_token":0.00000159984,"cache_read_input_token_cost":1.6968e-8}, + "cache_read_input_token_cost": 1.8354e-08, + "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"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"}],"input_cost_per_token":5.7684e-7,"output_cost_per_token":0.00000173052,"cache_read_input_token_cost":1.8354e-8}, "supports_audio_input": false, "supports_pdf_input": false, "supports_vision": false, @@ -44247,6 +44301,84 @@ "supports_system_messages": true, "supports_native_structured_output": true }, + "bedrock/ap-northeast-1/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.45e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/ap-south-1/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.41e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/ap-southeast-2/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.545e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.236e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/eu-west-1/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.41e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/eu-west-2/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 2.3e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.86e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/sa-east-1/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.45e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, "qwen.qwen3-vl-235b-a22b": { "input_cost_per_token": 5.3e-07, "litellm_provider": "bedrock_converse", @@ -46291,8 +46423,8 @@ "together_ai/zai-org/GLM-4.6": { "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", - "max_input_tokens": 200000, - "max_tokens": 200000, + "max_input_tokens": 202752, + "max_tokens": 202752, "metadata": { "successor": "together_ai/zai-org/GLM-5.2" }, @@ -46308,8 +46440,8 @@ "deprecation_date": "2026-04-02", "input_cost_per_token": 4.5e-07, "litellm_provider": "together_ai", - "max_input_tokens": 200000, - "max_tokens": 200000, + "max_input_tokens": 202752, + "max_tokens": 202752, "metadata": { "successor": "together_ai/zai-org/GLM-5.2" }, @@ -47158,7 +47290,7 @@ "mode": "chat", "output_cost_per_token": 1.2e-05, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "source": "https://aws.amazon.com/bedrock/pricing/", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -47192,7 +47324,7 @@ "mode": "chat", "output_cost_per_token": 3e-05, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "source": "https://aws.amazon.com/bedrock/pricing/", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -47225,7 +47357,7 @@ "mode": "chat", "output_cost_per_token": 3e-05, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "source": "https://aws.amazon.com/bedrock/pricing/", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -47276,7 +47408,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us-gov.nvidia.nemotron-nano-3-30b": { "input_cost_per_token": 7.2e-08, @@ -64215,6 +64347,25 @@ "supports_tool_choice": true, "supports_vision": false }, + "fireworks_ai/glm-5p3": { + "cache_read_input_token_cost": 2.6e-07, + "cache_read_input_token_cost_priority": 3.25e-07, + "input_cost_per_token": 1.4e-06, + "input_cost_per_token_priority": 1.75e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 4.4e-06, + "output_cost_per_token_priority": 5.5e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/accounts/fireworks/routers/glm-5p3-fast": { "cache_read_input_token_cost": 3.9e-07, "input_cost_per_token": 2.1e-06, @@ -64262,6 +64413,23 @@ "supports_tool_choice": true, "supports_vision": true }, + "fireworks_ai/glm-5p3-flash": { + "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_priority": 3.75e-08, + "input_cost_per_token": 1.5e-07, + "input_cost_per_token_priority": 1.875e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_priority": 6.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/accounts/fireworks/models/inkling": { "cache_read_input_token_cost": 1.7e-07, "input_cost_per_token": 1e-06, @@ -66550,13 +66718,13 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-flash-vision-exp": { - "input_cost_per_token": 2.156e-07, - "output_cost_per_token": 6.468e-07, - "cache_read_input_token_cost": 6.86e-09, + "input_cost_per_token": 2.2e-07, + "output_cost_per_token": 6.6e-07, + "cache_read_input_token_cost": 7e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -66592,8 +66760,8 @@ }, "openrouter/qwen/qwen3.8-27b": { "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.55e-06, - "cache_read_input_token_cost": 8.5e-08, + "output_cost_per_token": 2.5e-06, + "cache_read_input_token_cost": 5e-08, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 131072, @@ -66690,7 +66858,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "input_cost_per_token": 4e-08, - "output_cost_per_token": 8e-08, + "output_cost_per_token": 1.6e-07, "cache_read_input_token_cost": 1.6e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, @@ -67259,9 +67427,9 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "input_cost_per_token": 3.556e-08, - "output_cost_per_token": 7.112e-08, - "cache_read_input_token_cost": 7.112e-09, + "input_cost_per_token": 5.544e-08, + "output_cost_per_token": 1.1088e-07, + "cache_read_input_token_cost": 1.1088e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, @@ -68617,8 +68785,8 @@ "supports_web_search": true }, "openrouter/meta-llama/llama-4-maverick": { - "input_cost_per_token": 1.875e-07, - "output_cost_per_token": 6.525e-07, + "input_cost_per_token": 2e-07, + "output_cost_per_token": 8e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 16384, @@ -71300,8 +71468,8 @@ "input_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 384000, + "max_tokens": 384000, "mode": "chat", "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"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"}],"input_cost_per_token":1.5e-7,"output_cost_per_token":6e-7,"cache_read_input_token_cost":3e-9}, "output_cost_per_token": 1.2e-06, @@ -71317,14 +71485,15 @@ "supports_web_search": false }, "openrouter/~deepseek/deepseek-pro-latest": { - "cache_read_input_token_cost": 1.6968e-08, - "input_cost_per_token": 5.3328e-07, + "cache_read_input_token_cost": 1.8354e-08, + "input_cost_per_token": 5.7684e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.59984e-06, + "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"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"}],"input_cost_per_token":5.7684e-7,"output_cost_per_token":0.00000173052,"cache_read_input_token_cost":1.8354e-8}, + "output_cost_per_token": 1.73052e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -71344,7 +71513,7 @@ "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 8e-08, + "output_cost_per_token": 1.6e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -71573,8 +71742,8 @@ "input_cost_per_token": 9e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 3e-07, "source": "https://openrouter.ai/api/v1/models", @@ -72726,14 +72895,14 @@ "supports_web_search": false }, "openrouter/ibm-granite/granite-4.2-8b": { - "cache_read_input_token_cost": 1.5e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 117964, "max_tokens": 117964, "mode": "chat", - "output_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.5e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, diff --git a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py index b0d57cb6228..b0640e4f0dd 100644 --- a/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py +++ b/litellm/proxy/_experimental/mcp_server/auth/user_api_key_auth_mcp.py @@ -1154,7 +1154,7 @@ class MCPRequestHandler: Failures surface with the status the standard pipeline would give them, mirroring ``UserAPIKeyAuthExceptionHandler``: a disallowed route is the route gate's own 403, an - over-budget identity is a 429, a sub-check that raised its own ``HTTPException``/ + over-budget identity is a 422, a sub-check that raised its own ``HTTPException``/ ``ProxyException`` keeps that status, a transient database outage is a retryable 503, and only a genuinely unresolvable failure (a blocked team/project raises a bare ``Exception``, same as the standard pipeline's fallback) becomes the fail-closed 401. Collapsing every diff --git a/litellm/proxy/a2a/discovery.py b/litellm/proxy/a2a/discovery.py index e08c938f195..ff58c9c85ec 100644 --- a/litellm/proxy/a2a/discovery.py +++ b/litellm/proxy/a2a/discovery.py @@ -14,6 +14,7 @@ fetcher dispatches by ``discovery_mode``: pure-A2A fallback strategy returns 404 for these deployments. """ +from collections.abc import Mapping from enum import Enum from typing import Any, Final from urllib.parse import urlencode @@ -55,7 +56,7 @@ def _normalize_base_url(base_url: str) -> str: def _build_langgraph_platform_paths( - params: dict[str, Any] | None, + params: Mapping[str, object] | None, ) -> tuple[str, ...]: """Build the paths to try for LangGraph Platform discovery. @@ -71,7 +72,7 @@ def _build_langgraph_platform_paths( return tuple(f"{path}?{query}" for path in AGENT_CARD_WELL_KNOWN_PATHS) -def _paths_for_mode(mode: DiscoveryMode, params: dict[str, Any] | None) -> tuple[str, ...]: +def _paths_for_mode(mode: DiscoveryMode, params: Mapping[str, object] | None) -> tuple[str, ...]: if mode == DiscoveryMode.WELL_KNOWN_FALLBACK: return AGENT_CARD_WELL_KNOWN_PATHS if mode == DiscoveryMode.LANGGRAPH_PLATFORM: @@ -83,7 +84,7 @@ async def fetch_well_known_card( base_url: str, *, discovery_mode: DiscoveryMode = DiscoveryMode.WELL_KNOWN_FALLBACK, - params: dict[str, Any] | None = None, + params: Mapping[str, object] | None = None, timeout: float = DEFAULT_DISCOVERY_TIMEOUT_SECONDS, headers: dict[str, str] | None = None, ) -> dict[str, Any]: diff --git a/litellm/proxy/agent_endpoints/databricks_oauth.py b/litellm/proxy/agent_endpoints/databricks_oauth.py index 4c3b1bc084d..38a76ea6890 100644 --- a/litellm/proxy/agent_endpoints/databricks_oauth.py +++ b/litellm/proxy/agent_endpoints/databricks_oauth.py @@ -25,8 +25,9 @@ Config example:: import asyncio import base64 import hashlib +from collections.abc import Mapping from dataclasses import dataclass -from typing import Any, Final +from typing import Final import httpx @@ -43,7 +44,7 @@ _TOKEN_EXPIRY_BUFFER_SECONDS: Final = 60 _DEFAULT_TTL_SECONDS: Final = 3600 -def _resolve_secret(value: Any) -> str | None: +def _resolve_secret(value: object) -> str | None: """Resolve a config value, expanding ``os.environ/`` references.""" if not isinstance(value, str): return None @@ -75,7 +76,7 @@ class DatabricksAppOAuthConfig: def parse_databricks_oauth_config( - litellm_params: dict[str, Any] | None, + litellm_params: Mapping[str, object] | None, ) -> DatabricksAppOAuthConfig | None: """Build a Databricks App OAuth config from an agent's ``litellm_params``. @@ -191,7 +192,7 @@ class DatabricksAppOAuthTokenCache(InMemoryCache): except httpx.HTTPError as exc: raise ValueError(f"Databricks App OAuth token request failed: {exc}") from exc - body: Final = response.json() + body: Final[object] = response.json() if not isinstance(body, dict): raise ValueError( f"Databricks App OAuth token response returned non-object JSON (got {type(body).__name__})" @@ -215,7 +216,7 @@ databricks_app_oauth_token_cache: Final = DatabricksAppOAuthTokenCache() async def resolve_databricks_app_auth_header( - litellm_params: dict[str, Any] | None, + litellm_params: Mapping[str, object] | None, ) -> dict[str, str] | None: """Return ``{"Authorization": "Bearer "}`` for a Databricks App agent. diff --git a/litellm/proxy/auth/handle_jwt.py b/litellm/proxy/auth/handle_jwt.py index 803093ff93a..996911cdfaa 100644 --- a/litellm/proxy/auth/handle_jwt.py +++ b/litellm/proxy/auth/handle_jwt.py @@ -2468,7 +2468,22 @@ class JWTAuthManager: jwt_valid_token, handler, prisma_client, user_api_key_cache, parent_otel_span, proxy_logging_obj ) return {**admin_result, "user_object": identity.user_object} - return admin_result + if prisma_client is None: + return admin_result + try: + admin_user: Final = await get_user_object( + user_id=user_id, + user_email=user_email, + sso_user_id=user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + parent_otel_span=parent_otel_span, + proxy_logging_obj=proxy_logging_obj, + ) + except UserNotFoundError: + return admin_result + return {**admin_result, "user_object": admin_user} # Get team with model access ## Check if team_id is specified via x-litellm-team-id header diff --git a/litellm/proxy/auth/user_api_key_auth.py b/litellm/proxy/auth/user_api_key_auth.py index de0131772bc..4371ce4fda8 100644 --- a/litellm/proxy/auth/user_api_key_auth.py +++ b/litellm/proxy/auth/user_api_key_auth.py @@ -1714,6 +1714,16 @@ async def _user_api_key_auth_builder( jwt_claims = result.get("jwt_claims", None) agent_id: Final[str | None] = result.get("agent_id") + if ( + user_object is not None + and isinstance(user_object.metadata, dict) + and user_object.metadata.get("scim_active") is False + ): + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail=f"User={user_id} has been deactivated via SCIM. Keys owned by this user cannot be used.", + ) + if is_proxy_admin: # Proxy admins authenticate via auth_builder (full # access), not via a mapped virtual key. If diff --git a/litellm/proxy/client/cli/main.py b/litellm/proxy/client/cli/main.py index 63e38c93221..6d63acc7479 100644 --- a/litellm/proxy/client/cli/main.py +++ b/litellm/proxy/client/cli/main.py @@ -9,7 +9,15 @@ from litellm._version import version as litellm_version from litellm.proxy.client.health import HealthManagementClient from .commands.agents import agent_commands -from .commands.auth import auth_group, context_secret_vault, get_stored_api_key, login, logout, whoami +from .commands.auth import ( + CliContextObj, + auth_group, + context_secret_vault, + get_stored_api_key, + login, + logout, + whoami, +) from .commands.autoroute.commands import autoroute_group from .commands.chat import chat from .commands.config import config_commands, get_config_value, hidden_command_names @@ -126,7 +134,8 @@ def cli(ctx: click.Context, show_version: bool, base_url: str | None, api_key: s @click.pass_context def version(ctx: click.Context): """Show the LiteLLM Proxy CLI and server version.""" - print_version(ctx.obj.get("base_url"), ctx.obj.get("api_key")) + ctx_obj: Final[CliContextObj] = ctx.obj + print_version(ctx_obj.get("base_url"), ctx_obj.get("api_key")) # Add authentication commands as top-level commands diff --git a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py index d1576b68813..e3511d46544 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/generic_guardrail_api/__init__.py @@ -8,7 +8,7 @@ if TYPE_CHECKING: from litellm.types.guardrails import Guardrail, LitellmParams -def _get_config_value(litellm_params: Any, optional_params: Any, attribute_name: str) -> Any | None: +def _get_config_value(litellm_params: "LitellmParams", optional_params: object, attribute_name: str) -> Any | None: if optional_params is not None: value: Final = ( optional_params.get(attribute_name) diff --git a/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py b/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py index 7529c4ce3f3..1a6feb47215 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py +++ b/litellm/proxy/guardrails/guardrail_hooks/onyx/onyx.py @@ -6,7 +6,7 @@ # +-------------------------------------------------------------+ import os import uuid -from typing import TYPE_CHECKING, Any, Final, Literal, Optional +from typing import TYPE_CHECKING, Final, Literal, Optional import httpx from fastapi import HTTPException @@ -63,7 +63,7 @@ class OnyxGuardrail(CustomGuardrail): async def _validate_with_guard_server( self, - payload: Any, + payload: object, input_type: Literal["request", "response"], conversation_id: str, ) -> dict: diff --git a/litellm/proxy/hooks/responses_id_security.py b/litellm/proxy/hooks/responses_id_security.py index d9050489095..bdf7e2ab53d 100644 --- a/litellm/proxy/hooks/responses_id_security.py +++ b/litellm/proxy/hooks/responses_id_security.py @@ -40,7 +40,7 @@ _UNMANAGED_RESPONSE_ID_DETAIL: Final = ( _PROXY_ADMIN_ROLES: Final = frozenset({LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN.value}) -def _proxy_general_settings() -> Mapping[str, Any]: +def _proxy_general_settings() -> Mapping[str, object]: from litellm.proxy.proxy_server import general_settings return general_settings @@ -107,7 +107,7 @@ def _is_responses_api_create_route(request_route: str | None) -> bool: class ResponsesIDSecurity(CustomLogger): def __init__( self, - general_settings_reader: Callable[[], Mapping[str, Any]] = _proxy_general_settings, + general_settings_reader: Callable[[], Mapping[str, object]] = _proxy_general_settings, signing_key_reader: Callable[[], str | None] = _proxy_signing_key, ) -> None: self._general_settings_reader: Final = general_settings_reader @@ -307,7 +307,7 @@ class ResponsesIDSecurity(CustomLogger): data: dict, user_api_key_dict: "UserAPIKeyAuth", response: LLMResponseTypes, - ) -> Any: + ) -> LLMResponseTypes: """ Queue response IDs for batch processing instead of writing directly to DB. diff --git a/litellm/proxy/logging_endpoints/callback_logs_endpoints.py b/litellm/proxy/logging_endpoints/callback_logs_endpoints.py index cecadc03d71..4a1079871b0 100644 --- a/litellm/proxy/logging_endpoints/callback_logs_endpoints.py +++ b/litellm/proxy/logging_endpoints/callback_logs_endpoints.py @@ -15,6 +15,7 @@ self-describing `StandardLoggingPayload`, so completions/responses can use it to """ import uuid +from collections.abc import Mapping from datetime import datetime, timezone from typing import Any, Final @@ -48,7 +49,7 @@ class CallbackLogsReplayer: """ @staticmethod - def _epoch_to_datetime(value: Any) -> datetime: + def _epoch_to_datetime(value: object) -> datetime: """`StandardLoggingPayload` stores startTime/endTime as float epoch seconds.""" if isinstance(value, (int, float)): return datetime.fromtimestamp(float(value), tz=timezone.utc) @@ -114,7 +115,7 @@ class CallbackLogsReplayer: return logging_obj @staticmethod - def _response_obj_from_payload(payload: dict[str, Any]) -> dict[str, Any]: + def _response_obj_from_payload(payload: Mapping[str, object]) -> dict[str, object]: """Minimal response object so usage-derived spend-log fields resolve.""" return { "id": payload.get("id"), diff --git a/litellm/proxy/management_endpoints/internal_user_endpoints.py b/litellm/proxy/management_endpoints/internal_user_endpoints.py index 4832c2f4c21..029a968e156 100644 --- a/litellm/proxy/management_endpoints/internal_user_endpoints.py +++ b/litellm/proxy/management_endpoints/internal_user_endpoints.py @@ -1571,7 +1571,7 @@ async def _update_single_user_helper( await _invalidate_user_spend_counter_if_changed(non_default_values) - if "model_max_budget" in non_default_values: + if "model_max_budget" in non_default_values or "metadata" in data_json: await evict_and_broadcast( cache_keys=(non_default_values["user_id"],), user_api_key_cache=user_api_key_cache, diff --git a/litellm/proxy/management_endpoints/management_v1/spend_logs.py b/litellm/proxy/management_endpoints/management_v1/spend_logs.py index f6907a7f87a..1cbc454ca5e 100644 --- a/litellm/proxy/management_endpoints/management_v1/spend_logs.py +++ b/litellm/proxy/management_endpoints/management_v1/spend_logs.py @@ -1,7 +1,7 @@ """`/management/v1/spend_logs` facets.""" from datetime import datetime, timezone -from typing import Annotated, Any, Final, Literal +from typing import Annotated, Final, Literal from fastapi import APIRouter, Depends, Query, Request @@ -39,7 +39,7 @@ async def _spend_log_scope_clause( user_api_key_dict: UserAPIKeyAuth, prisma_client: PrismaClient, next_param_index: int, -) -> tuple[str | None, tuple[Any, ...]]: +) -> tuple[str | None, tuple[str | list[str], ...]]: """SQL predicate restricting the facet to spend logs this caller may read. Returns ``(None, ())`` for a proxy admin. Mirrors the scoping ``/spend/logs/ui`` @@ -101,8 +101,8 @@ async def _list_spend_log_facet( ) column_sql: Final = "end_user" if column == "end_user" else '"user"' - window_params: Final[tuple[Any, ...]] = (_as_utc(start_time), _as_utc(end_time)) - search_params: Final[tuple[Any, ...]] = (f"%{escape_like(q)}%",) if q else () + window_params: Final[tuple[datetime, datetime]] = (_as_utc(start_time), _as_utc(end_time)) + search_params: Final[tuple[str, ...]] = (f"%{escape_like(q)}%",) if q else () search_clause: Final = (f"{column_sql} ILIKE ${len(window_params) + 1} ESCAPE '\\'",) if q else () scope_clause, scope_params = await _spend_log_scope_clause( diff --git a/litellm/proxy/management_endpoints/scim/scim_v2.py b/litellm/proxy/management_endpoints/scim/scim_v2.py index 2b74dc1e838..b676c0ddb82 100644 --- a/litellm/proxy/management_endpoints/scim/scim_v2.py +++ b/litellm/proxy/management_endpoints/scim/scim_v2.py @@ -46,6 +46,7 @@ from litellm.proxy._types import ( ) from litellm.proxy.auth.auth_checks import _delete_cache_key_object from litellm.proxy.auth.user_api_key_auth import user_api_key_auth +from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import evict_and_broadcast from litellm.proxy.common_utils.http_parsing_utils import _safe_get_request_headers from litellm.proxy.management_endpoints.internal_user_endpoints import new_user from litellm.proxy.management_endpoints.scim.scim_transformations import ( @@ -1804,6 +1805,9 @@ async def update_user( where={"user_id": user_id}, data=update_data, ) + from litellm.proxy.proxy_server import user_api_key_cache + + await evict_and_broadcast(cache_keys=(user_id,), user_api_key_cache=user_api_key_cache) if client_set_active: new_active: Final = _scim_active_value(metadata) @@ -2375,6 +2379,9 @@ async def patch_user( where={"user_id": user_id}, data=update_data, ) + from litellm.proxy.proxy_server import user_api_key_cache + + await evict_and_broadcast(cache_keys=(user_id,), user_api_key_cache=user_api_key_cache) if new_active is not None and new_active != (True if prev_active is None else prev_active): await _set_user_keys_blocked(user_id=user_id, blocked=not new_active) diff --git a/litellm/proxy/management_helpers/object_permission_utils.py b/litellm/proxy/management_helpers/object_permission_utils.py index daab38d3662..437e6763502 100644 --- a/litellm/proxy/management_helpers/object_permission_utils.py +++ b/litellm/proxy/management_helpers/object_permission_utils.py @@ -8,7 +8,7 @@ from collections.abc import Mapping, Sequence from collections.abc import Set as AbstractSet from dataclasses import dataclass from types import MappingProxyType -from typing import TYPE_CHECKING, Any, Final, Optional +from typing import TYPE_CHECKING, Final, Optional from fastapi import HTTPException, status from pydantic import TypeAdapter @@ -230,7 +230,7 @@ def _dedupe_preserving_order(values: list[str]) -> list[str]: return result -def _mcp_server_identifier_matches(server: Any, identifier: str) -> bool: +def _mcp_server_identifier_matches(server: object, identifier: str) -> bool: return identifier in { getattr(server, "server_id", None), getattr(server, "alias", None), diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py index a95ee87fd31..d97ddb9a909 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/gemini_passthrough_logging_handler.py @@ -147,7 +147,7 @@ class GeminiPassthroughLoggingHandler: - Creates standard logging object - Logs in litellm callbacks """ - kwargs: dict[str, Any] = {} + kwargs: dict[str, object] = {} model = model or GeminiPassthroughLoggingHandler.extract_model_from_url(url_route) complete_streaming_response: Final = GeminiPassthroughLoggingHandler._build_complete_streaming_response( all_chunks=all_chunks, diff --git a/litellm/proxy/public_endpoints/public_endpoints.py b/litellm/proxy/public_endpoints/public_endpoints.py index e395f56194f..26a5c44fce1 100644 --- a/litellm/proxy/public_endpoints/public_endpoints.py +++ b/litellm/proxy/public_endpoints/public_endpoints.py @@ -199,9 +199,13 @@ def _build_endpoints(raw: _ProvidersFile) -> list[_EndpointEntry]: return result +_PROVIDERS_FILE_ADAPTER: Final = TypeAdapter(_ProvidersFile) +_PROVIDER_CREATE_FIELDS_ADAPTER: Final = TypeAdapter(list[ProviderCreateInfo]) + + def _load_endpoints() -> list[_EndpointEntry]: - raw: Final[_ProvidersFile] = json.loads( - files("litellm").joinpath("provider_endpoints_support_backup.json").read_text(encoding="utf-8") + raw: Final = _PROVIDERS_FILE_ADAPTER.validate_python( + json.loads(files("litellm").joinpath("provider_endpoints_support_backup.json").read_text(encoding="utf-8")) ) return _build_endpoints(raw) @@ -398,7 +402,7 @@ async def get_provider_fields() -> list[ProviderCreateInfo]: ) with open(provider_create_fields_path, "r") as f: - provider_create_fields: Final = json.load(f) + provider_create_fields: Final = _PROVIDER_CREATE_FIELDS_ADAPTER.validate_python(json.load(f)) return provider_create_fields diff --git a/litellm/responses/litellm_completion_transformation/transformation.py b/litellm/responses/litellm_completion_transformation/transformation.py index cf3075ee28d..3ca2cc28c9a 100644 --- a/litellm/responses/litellm_completion_transformation/transformation.py +++ b/litellm/responses/litellm_completion_transformation/transformation.py @@ -454,6 +454,7 @@ class LiteLLMCompletionResponsesConfig: "stream": stream, "metadata": kwargs.get("metadata"), "service_tier": kwargs.get("service_tier"), + "safety_identifier": responses_api_request.get("safety_identifier"), "web_search_options": web_search_options, "response_format": response_format, "reasoning_effort": reasoning.effort, diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index c0a06364261..05a6df6d5af 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -100,6 +100,8 @@ class TokenCounter: def from_cl100k_ranks(rank_file: str) -> TokenCounter: ... @staticmethod def from_o200k_ranks(rank_file: str) -> TokenCounter: ... + @staticmethod + def from_tiktoken(encoding: str) -> TokenCounter: ... def acount_request(self, body: bytes) -> Future[dict[str, object]]: ... def gil_stats() -> dict[str, int]: ... diff --git a/litellm/types/llms/anthropic.py b/litellm/types/llms/anthropic.py index bcd24695f25..c59c88698f7 100644 --- a/litellm/types/llms/anthropic.py +++ b/litellm/types/llms/anthropic.py @@ -411,6 +411,7 @@ class AnthropicMessagesRequestOptionalParams(TypedDict, total=False): output_config: AnthropicOutputConfig | None # Configuration for Claude's output behavior cache_control: dict[str, Any] | None # Automatic prompt caching reasoning_effort: str | None + safeguards: ReadOnly[list[dict[str, object]] | None] class AnthropicMessagesRequest(AnthropicMessagesRequestOptionalParams, total=False): @@ -530,6 +531,7 @@ class AnthropicStopDetails(TypedDict, total=False): class MessageDelta(TypedDict, total=False): stop_reason: str | None stop_details: ReadOnly[AnthropicStopDetails] + safeguard_results: ReadOnly[list[dict[str, object]]] class ServerToolUsage(TypedDict, total=False): @@ -600,6 +602,7 @@ class MessageChunk(TypedDict, total=False): stop_reason: str | None stop_sequence: str | None usage: UsageDelta + safeguard_results: ReadOnly[list[dict[str, object]]] class MessageStartBlock(TypedDict): diff --git a/litellm/types/llms/anthropic_messages/anthropic_response.py b/litellm/types/llms/anthropic_messages/anthropic_response.py index 038a23a3ca2..1d4c3cdc864 100644 --- a/litellm/types/llms/anthropic_messages/anthropic_response.py +++ b/litellm/types/llms/anthropic_messages/anthropic_response.py @@ -97,3 +97,4 @@ class AnthropicMessagesResponse(TypedDict, total=False): type: Literal["message"] | None usage: AnthropicUsage | None context_management: NotRequired[ContextManagementResponse] + safeguard_results: NotRequired[ReadOnly[list[dict[str, object]]]] diff --git a/litellm/types/utils.py b/litellm/types/utils.py index cc9b7931f56..5a80644347e 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -255,6 +255,7 @@ class ModelInfoBase(ProviderSpecificModelInfo, total=False): cache_creation_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing cache_read_input_token_cost: float | None cache_read_input_audio_token_cost: ReadOnly[float | None] + cache_read_input_image_token_cost: ReadOnly[float | None] cache_read_input_token_cost_flex: float | None # OpenAI flex service tier pricing cache_read_input_token_cost_priority: float | None # OpenAI priority service tier pricing cache_read_input_token_cost_ultrafast: ReadOnly[float | None] # OpenAI ultrafast service tier pricing @@ -3635,6 +3636,7 @@ class CustomPricingLiteLLMParams(MirroredPricingParams): cache_read_input_token_cost_above_272k_tokens_priority: float | None = None cache_read_input_token_cost_above_272k_tokens_flex: float | None = None cache_read_input_audio_token_cost: float | None = None + cache_read_input_image_token_cost: float | None = None input_cost_per_character_above_128k_tokens: float | None = None input_cost_per_audio_token: float | None = None input_cost_per_token_cache_hit: float | None = None diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index bf65675f3c5..5b686f0a4b5 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -1327,7 +1327,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 2048, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "anthropic.claude-mythos-preview": { "input_cost_per_token": 0, @@ -1381,7 +1381,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 2048, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-opus-4-7": { "bedrock_converse_supports_strict_tools": false, @@ -1419,7 +1419,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 2048, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-opus-4-7": { "bedrock_converse_supports_strict_tools": false, @@ -1531,7 +1531,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "anthropic.claude-fable-5-1": { "cache_creation_input_token_cost": 1.25e-05, @@ -1570,7 +1570,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-fable-5": { "cache_creation_input_token_cost": 1.25e-05, @@ -1608,7 +1608,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-fable-5-1": { "cache_creation_input_token_cost": 1.25e-05, @@ -1647,7 +1647,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-fable-5": { "cache_creation_input_token_cost": 1.375e-05, @@ -1685,7 +1685,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-fable-5-1": { "cache_creation_input_token_cost": 1.375e-05, @@ -1724,7 +1724,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-fable-5": { "cache_creation_input_token_cost": 1.375e-05, @@ -1837,7 +1837,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-opus-5": { "bedrock_converse_supports_strict_tools": false, @@ -1875,7 +1875,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-opus-5": { "bedrock_converse_supports_strict_tools": false, @@ -1913,7 +1913,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-opus-5": { "bedrock_converse_supports_strict_tools": false, @@ -2063,7 +2063,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, @@ -2102,7 +2102,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, @@ -2141,7 +2141,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-opus-4-8": { "bedrock_converse_supports_strict_tools": false, @@ -2329,7 +2329,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-sonnet-5": { "bedrock_converse_supports_strict_tools": false, @@ -2368,7 +2368,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-sonnet-5": { "bedrock_converse_supports_strict_tools": false, @@ -2407,7 +2407,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-sonnet-5": { "bedrock_converse_supports_strict_tools": false, @@ -2556,7 +2556,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "global.anthropic.claude-sonnet-4-6": { "supports_adaptive_thinking": true, @@ -2591,7 +2591,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us.anthropic.claude-sonnet-4-6": { "supports_adaptive_thinking": true, @@ -2626,7 +2626,7 @@ "supports_output_config": true, "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "eu.anthropic.claude-sonnet-4-6": { "supports_adaptive_thinking": true, @@ -3740,6 +3740,21 @@ "supports_vision": true, "supports_web_search": true }, + "azure_ai/gpt-image-2": { + "cache_read_input_image_token_cost": 2e-06, + "cache_read_input_token_cost": 1.25e-06, + "input_cost_per_image_token": 8e-06, + "input_cost_per_token": 5e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image_token": 3e-05, + "source": "https://azure.microsoft.com/en-us/pricing/details/cognitive-services/openai-service/", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ], + "supports_vision": true + }, "azure_ai/codex-mini": { "cache_read_input_token_cost": 3.75e-07, "deprecation_date": "2026-11-15", @@ -11159,6 +11174,20 @@ ], "deprecation_date": "2026-10-01" }, + "azure_ai/MAI-Image-2.5-Pro": { + "deprecation_date": "2026-10-01", + "input_cost_per_image_token": 8e-06, + "input_cost_per_token": 5e-06, + "litellm_provider": "azure_ai", + "mode": "image_generation", + "output_cost_per_image": 0.1085, + "output_cost_per_image_token": 0.000106, + "source": "https://techcommunity.microsoft.com/blog/azure-ai-foundry-blog/introducing-mai-image-2-5-pro-and-mai-voice-2-flash-in-microsoft-foundry/4539446", + "supported_endpoints": [ + "/v1/images/generations", + "/v1/images/edits" + ] + }, "azure_ai/MAI-Image-2e": { "deprecation_date": "2026-08-15", "input_cost_per_token": 5e-06, @@ -21888,6 +21917,7 @@ "supports_tool_choice": true }, "deepseek/deepseek-coder": { + "cache_read_input_token_cost": 1.4e-08, "input_cost_per_token": 1.4e-07, "input_cost_per_token_cache_hit": 1.4e-08, "litellm_provider": "deepseek", @@ -21902,6 +21932,7 @@ "supports_tool_choice": true }, "deepseek/deepseek-r1": { + "cache_read_input_token_cost": 1.4e-07, "input_cost_per_token": 5.5e-07, "input_cost_per_token_cache_hit": 1.4e-07, "litellm_provider": "deepseek", @@ -21957,6 +21988,7 @@ "supports_tool_choice": true }, "deepseek/deepseek-v3.2": { + "cache_read_input_token_cost": 2.8e-08, "input_cost_per_token": 2.8e-07, "input_cost_per_token_cache_hit": 2.8e-08, "litellm_provider": "deepseek", @@ -23796,6 +23828,25 @@ "supports_tool_choice": true, "supports_vision": false }, + "fireworks_ai/deepseek-v4-pro-0813": { + "cache_read_input_token_cost": 4.4e-08, + "cache_read_input_token_cost_priority": 5.5e-08, + "input_cost_per_token": 1.32e-06, + "input_cost_per_token_priority": 1.65e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "output_cost_per_token": 3.96e-06, + "output_cost_per_token_priority": 4.95e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/accounts/fireworks/models/firefunction-v2": { "input_cost_per_token": 9e-07, "litellm_provider": "fireworks_ai", @@ -24182,7 +24233,7 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": false }, "fireworks_ai/accounts/fireworks/models/mixtral-8x22b-instruct-hf": { "input_cost_per_token": 1.2e-06, @@ -24508,7 +24559,7 @@ "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": false }, "fireworks_ai/qwen3p7-plus": { "cache_read_input_token_cost": 8e-08, @@ -35244,6 +35295,7 @@ "mode": "chat", "output_cost_per_token": 3e-06, "source": "https://console.groq.com/docs/model/qwen/qwen3.6-27b", + "deprecation_date": "2026-09-14", "supports_function_calling": true, "supports_reasoning": true, "supports_response_schema": false, @@ -41472,6 +41524,7 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v3.2-exp": { + "cache_read_input_token_cost": 2e-08, "deprecation_date": "2026-09-28", "input_cost_per_token": 2.7e-07, "input_cost_per_token_cache_hit": 2e-08, @@ -41494,6 +41547,7 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-r1": { + "cache_read_input_token_cost": 1.4e-07, "input_cost_per_token": 7e-07, "input_cost_per_token_cache_hit": 1.4e-07, "litellm_provider": "openrouter", @@ -41537,21 +41591,21 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 4.22298e-07, + "input_cost_per_token": 9.42906e-07, "input_cost_per_token_cache_hit": 4.4e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, "max_tokens": 384000, "mode": "chat", - "output_cost_per_token": 8.44596e-07, + "output_cost_per_token": 1.885812e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 3.51915e-08, + "cache_read_input_token_cost": 7.85755e-08, "supports_audio_input": false, "supports_pdf_input": false, "supports_vision": false, @@ -41579,22 +41633,22 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "input_cost_per_token": 5.3328e-07, - "input_cost_per_token_cache_hit": 4.4e-08, + "input_cost_per_token": 5.7684e-07, + "input_cost_per_token_cache_hit": 1.9272e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.59984e-06, + "output_cost_per_token": 1.73052e-06, "source": "https://openrouter.ai/api/v1/models", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, - "cache_read_input_token_cost": 1.6968e-08, - "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"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"}],"input_cost_per_token":5.3328e-7,"output_cost_per_token":0.00000159984,"cache_read_input_token_cost":1.6968e-8}, + "cache_read_input_token_cost": 1.8354e-08, + "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"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"}],"input_cost_per_token":5.7684e-7,"output_cost_per_token":0.00000173052,"cache_read_input_token_cost":1.8354e-8}, "supports_audio_input": false, "supports_pdf_input": false, "supports_vision": false, @@ -44247,6 +44301,84 @@ "supports_system_messages": true, "supports_native_structured_output": true }, + "bedrock/ap-northeast-1/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.45e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/ap-south-1/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.41e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/ap-southeast-2/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.545e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.236e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/eu-west-1/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.41e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/eu-west-2/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 2.3e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.86e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, + "bedrock/sa-east-1/qwen.qwen3-next-80b-a3b": { + "input_cost_per_token": 1.8e-07, + "litellm_provider": "bedrock", + "max_input_tokens": 128000, + "max_output_tokens": 8192, + "max_tokens": 8192, + "mode": "chat", + "output_cost_per_token": 1.45e-06, + "source": "https://aws.amazon.com/bedrock/pricing/", + "supports_function_calling": true, + "supports_native_structured_output": true, + "supports_system_messages": true + }, "qwen.qwen3-vl-235b-a22b": { "input_cost_per_token": 5.3e-07, "litellm_provider": "bedrock_converse", @@ -46291,8 +46423,8 @@ "together_ai/zai-org/GLM-4.6": { "input_cost_per_token": 6e-07, "litellm_provider": "together_ai", - "max_input_tokens": 200000, - "max_tokens": 200000, + "max_input_tokens": 202752, + "max_tokens": 202752, "metadata": { "successor": "together_ai/zai-org/GLM-5.2" }, @@ -46308,8 +46440,8 @@ "deprecation_date": "2026-04-02", "input_cost_per_token": 4.5e-07, "litellm_provider": "together_ai", - "max_input_tokens": 200000, - "max_tokens": 200000, + "max_input_tokens": 202752, + "max_tokens": 202752, "metadata": { "successor": "together_ai/zai-org/GLM-5.2" }, @@ -47158,7 +47290,7 @@ "mode": "chat", "output_cost_per_token": 1.2e-05, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "source": "https://aws.amazon.com/bedrock/pricing/", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -47192,7 +47324,7 @@ "mode": "chat", "output_cost_per_token": 3e-05, "prompt_cache_min_tokens": 1024, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "source": "https://aws.amazon.com/bedrock/pricing/", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -47225,7 +47357,7 @@ "mode": "chat", "output_cost_per_token": 3e-05, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json", + "source": "https://aws.amazon.com/bedrock/pricing/", "supports_adaptive_thinking": true, "supports_assistant_prefill": false, "supports_computer_use": true, @@ -47276,7 +47408,7 @@ "bedrock_output_config_effort_ceiling": "xhigh", "supports_parallel_tool_use_config": true, "prompt_cache_min_tokens": 512, - "source": "https://pricing.us-east-1.amazonaws.com/offers/v1.0/aws/AmazonBedrockFoundationModels/current/index.json" + "source": "https://aws.amazon.com/bedrock/pricing/" }, "us-gov.nvidia.nemotron-nano-3-30b": { "input_cost_per_token": 7.2e-08, @@ -64215,6 +64347,25 @@ "supports_tool_choice": true, "supports_vision": false }, + "fireworks_ai/glm-5p3": { + "cache_read_input_token_cost": 2.6e-07, + "cache_read_input_token_cost_priority": 3.25e-07, + "input_cost_per_token": 1.4e-06, + "input_cost_per_token_priority": 1.75e-06, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 4.4e-06, + "output_cost_per_token_priority": 5.5e-06, + "source": "https://api.fireworks.ai/v1/serverless/models", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": false + }, "fireworks_ai/accounts/fireworks/routers/glm-5p3-fast": { "cache_read_input_token_cost": 3.9e-07, "input_cost_per_token": 2.1e-06, @@ -64262,6 +64413,23 @@ "supports_tool_choice": true, "supports_vision": true }, + "fireworks_ai/glm-5p3-flash": { + "cache_read_input_token_cost": 3e-08, + "cache_read_input_token_cost_priority": 3.75e-08, + "input_cost_per_token": 1.5e-07, + "input_cost_per_token_priority": 1.875e-07, + "litellm_provider": "fireworks_ai", + "max_input_tokens": 1048576, + "max_tokens": 1048576, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_priority": 6.25e-07, + "source": "https://api.fireworks.ai/v1/serverless/models", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "fireworks_ai/accounts/fireworks/models/inkling": { "cache_read_input_token_cost": 1.7e-07, "input_cost_per_token": 1e-06, @@ -66550,13 +66718,13 @@ "supports_web_search": false }, "openrouter/deepseek/deepseek-v4-flash-vision-exp": { - "input_cost_per_token": 2.156e-07, - "output_cost_per_token": 6.468e-07, - "cache_read_input_token_cost": 6.86e-09, + "input_cost_per_token": 2.2e-07, + "output_cost_per_token": 6.6e-07, + "cache_read_input_token_cost": 7e-09, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 262144, - "max_tokens": 262144, + "max_output_tokens": 943718, + "max_tokens": 943718, "mode": "chat", "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, @@ -66592,8 +66760,8 @@ }, "openrouter/qwen/qwen3.8-27b": { "input_cost_per_token": 2e-07, - "output_cost_per_token": 2.55e-06, - "cache_read_input_token_cost": 8.5e-08, + "output_cost_per_token": 2.5e-06, + "cache_read_input_token_cost": 5e-08, "litellm_provider": "openrouter", "max_input_tokens": 1000000, "max_output_tokens": 131072, @@ -66690,7 +66858,7 @@ }, "openrouter/deepseek/deepseek-v4-flash-0731": { "input_cost_per_token": 4e-08, - "output_cost_per_token": 8e-08, + "output_cost_per_token": 1.6e-07, "cache_read_input_token_cost": 1.6e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, @@ -67259,9 +67427,9 @@ "supports_web_search": true }, "openrouter/deepseek/deepseek-v4-flash": { - "input_cost_per_token": 3.556e-08, - "output_cost_per_token": 7.112e-08, - "cache_read_input_token_cost": 7.112e-09, + "input_cost_per_token": 5.544e-08, + "output_cost_per_token": 1.1088e-07, + "cache_read_input_token_cost": 1.1088e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, @@ -68617,8 +68785,8 @@ "supports_web_search": true }, "openrouter/meta-llama/llama-4-maverick": { - "input_cost_per_token": 1.875e-07, - "output_cost_per_token": 6.525e-07, + "input_cost_per_token": 2e-07, + "output_cost_per_token": 8e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 16384, @@ -71300,8 +71468,8 @@ "input_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 384000, + "max_tokens": 384000, "mode": "chat", "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"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"}],"input_cost_per_token":1.5e-7,"output_cost_per_token":6e-7,"cache_read_input_token_cost":3e-9}, "output_cost_per_token": 1.2e-06, @@ -71317,14 +71485,15 @@ "supports_web_search": false }, "openrouter/~deepseek/deepseek-pro-latest": { - "cache_read_input_token_cost": 1.6968e-08, - "input_cost_per_token": 5.3328e-07, + "cache_read_input_token_cost": 1.8354e-08, + "input_cost_per_token": 5.7684e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.59984e-06, + "off_peak_pricing": {"windows":[{"weekdays":["saturday","sunday"],"hours_utc":"00:00-00:00"},{"weekdays":["monday","tuesday","wednesday","thursday","friday"],"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"}],"input_cost_per_token":5.7684e-7,"output_cost_per_token":0.00000173052,"cache_read_input_token_cost":1.8354e-8}, + "output_cost_per_token": 1.73052e-06, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -71344,7 +71513,7 @@ "max_output_tokens": 943718, "max_tokens": 943718, "mode": "chat", - "output_cost_per_token": 8e-08, + "output_cost_per_token": 1.6e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, @@ -71573,8 +71742,8 @@ "input_cost_per_token": 9e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, - "max_output_tokens": 943718, - "max_tokens": 943718, + "max_output_tokens": 131072, + "max_tokens": 131072, "mode": "chat", "output_cost_per_token": 3e-07, "source": "https://openrouter.ai/api/v1/models", @@ -72726,14 +72895,14 @@ "supports_web_search": false }, "openrouter/ibm-granite/granite-4.2-8b": { - "cache_read_input_token_cost": 1.5e-08, - "input_cost_per_token": 6e-08, + "cache_read_input_token_cost": 5e-08, + "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 117964, "max_tokens": 117964, "mode": "chat", - "output_cost_per_token": 2.5e-07, + "output_cost_per_token": 1.5e-07, "source": "https://openrouter.ai/api/v1/models", "supports_audio_input": false, "supports_function_calling": true, diff --git a/model_prices_and_context_window.schema.json b/model_prices_and_context_window.schema.json index 509f957b8d1..5b0a23adfea 100644 --- a/model_prices_and_context_window.schema.json +++ b/model_prices_and_context_window.schema.json @@ -137,6 +137,10 @@ "type": "number", "minimum": 0 }, + "cache_read_input_image_token_cost": { + "type": "number", + "minimum": 0 + }, "cache_read_input_token_cost": { "type": "number", "minimum": 0, diff --git a/tests/base_sdk_tests/check_base_sdk_install.py b/tests/base_sdk_tests/check_base_sdk_install.py index 190a900faf9..f680ba92645 100644 --- a/tests/base_sdk_tests/check_base_sdk_install.py +++ b/tests/base_sdk_tests/check_base_sdk_install.py @@ -50,6 +50,17 @@ def check_completion() -> str: return "mock completion round-trips" +def check_mcp_install_guidance() -> str: + try: + import litellm.experimental_mcp_client + except ImportError as error: + _require("pip install 'litellm[mcp]'" in str(error), f"missing MCP installation guidance: {error}") + _require(isinstance(error.__cause__, ModuleNotFoundError), "original missing-dependency cause was lost") + _require(error.__cause__.name == "mcp", f"unexpected missing dependency: {error.__cause__}") + return "optional MCP client explains how to install litellm[mcp]" + raise AssertionError("MCP client imported without the MCP extra") + + def check_embedding() -> str: import litellm @@ -109,6 +120,7 @@ def check_bedrock_credential_resolution() -> str: CHECKS: tuple[tuple[str, Callable[[], str]], ...] = ( ("environment is base-only", check_environment_is_base_only), ("import litellm", check_import), + ("optional MCP installation guidance", check_mcp_install_guidance), ("chat completion", check_completion), ("embedding", check_embedding), ("bundled model metadata", check_bundled_model_metadata), diff --git a/tests/e2e/AGENTS.md b/tests/e2e/AGENTS.md index 78ba2ec4ed1..c25e958242f 100644 --- a/tests/e2e/AGENTS.md +++ b/tests/e2e/AGENTS.md @@ -122,7 +122,7 @@ E2E_FIXTURE_MODE=replay E2E_FIXTURE_DIR=/tmp/e2e-fixtures E2E_RESET_SPEND_LOGS=1 Point the proxy at bogus provider credentials for the replay run and it still has to pass: that is the whole proof that nothing left the process. Bundles are never committed. `tests/e2e/.fixtures` is gitignored because a bundle holds verbatim provider response bodies and hard-fails after seven days. CI records and replays this lane on a schedule in `.github/workflows/e2e_record_replay.yml`, publishing the bundle as a private `e2e-fixtures-bundle` artifact instead of committing it, selecting the tests with the `@pytest.mark.replayable` marker, and proving the bogus-credentials replay hermetic by counting provider egress with `.github/scripts/e2e_egress_sentinel.py` -Current limits: Bedrock cannot be mounted (SigV4 signs the Host header, so a rewritten api_base fails signature verification), deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base), and a file upload routed by `custom_llm_provider` through the proxy's `files_settings` block never passes a deployment at all, so the batches `model_param` and `provider_fallback` scenarios keep uploading live in every mode +Current limits: Bedrock cannot be mounted in record or replay (SigV4 signs the Host header, so a rewritten api_base fails signature verification); a test that needs to observe the Converse body registers its own `LiveEdge` with `provider_edge_bedrock.bedrock_signer` re-signing the forwarded request, and carries the `provider_edge_host` opt-in marker because the gateway must reach the pytest host, which the Buildkite ephemeral stack cannot (the GitHub changed-e2e lane, whose gateways run on the runner, sets `E2E_PROVIDER_EDGE_HOST_REACHABLE`). Deployments baked into the proxy's config file cannot be edge-wired (only `/model/new` registrations can carry the edge api_base), and a file upload routed by `custom_llm_provider` through the proxy's `files_settings` block never passes a deployment at all, so the batches `model_param` and `provider_fallback` scenarios keep uploading live in every mode ## Typing diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 52a634693c5..8776d00d502 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -31,6 +31,7 @@ from e2e_config import ( MANAGED_FILES_OPT_IN_ENV, MCP_OAUTH_LIVE_OPT_IN_ENV, PROMPT_CACHING_OPT_IN_ENV, + PROVIDER_EDGE_HOST_OPT_IN_ENV, PROXY_BASE_URL, REDIS_CHAOS_OPT_IN_ENV, WEEKLY_ANOMALY_OPT_IN_ENV, @@ -59,6 +60,7 @@ OPT_IN_MARKERS: Final = MappingProxyType( "redis_chaos": REDIS_CHAOS_OPT_IN_ENV, "cli_determinism": CLI_DETERMINISM_OPT_IN_ENV, "mcp_oauth_live": MCP_OAUTH_LIVE_OPT_IN_ENV, + "provider_edge_host": PROVIDER_EDGE_HOST_OPT_IN_ENV, } ) @@ -143,6 +145,11 @@ def pytest_configure(config: pytest.Config) -> None: "mcp_oauth_live: real Linear OAuth consent via a captured browser session; deselected unless " "E2E_MCP_OAUTH_LIVE is set", ) + config.addinivalue_line( + "markers", + "provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the " + "gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set", + ) def pytest_sessionstart(session: pytest.Session) -> None: diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index a79c158f9c4..14de4619664 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -146,6 +146,7 @@ PROMPT_CACHING_OPT_IN_ENV = "E2E_PROMPT_CACHING_STACK" REDIS_CHAOS_OPT_IN_ENV = "E2E_REDIS_CHAOS" CLI_DETERMINISM_OPT_IN_ENV = "E2E_CLI_DETERMINISM" MCP_OAUTH_LIVE_OPT_IN_ENV: Final = "E2E_MCP_OAUTH_LIVE" +PROVIDER_EDGE_HOST_OPT_IN_ENV: Final = "E2E_PROVIDER_EDGE_HOST_REACHABLE" ANOMALY_SESSIONS = int(os.environ.get("E2E_ANOMALY_SESSIONS", "6")) ANOMALY_TURNS_PER_SESSION = int(os.environ.get("E2E_ANOMALY_TURNS_PER_SESSION", "6")) ANOMALY_TURN_ATTEMPTS = int(os.environ.get("E2E_ANOMALY_TURN_ATTEMPTS", "3")) diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index 4184b6cbefc..d4978601b20 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -95,7 +95,7 @@ class UnauthorizedError(BaseModel): class RateLimitedError(BaseModel): kind: Literal["rate_limited"] = "rate_limited" retry_after_seconds: int | None = None - # litellm overloads 429 for budget_exceeded too, so keep the body to tell them apart. + # keep the body so callers can tell limiter kinds apart. body: str = "" diff --git a/tests/e2e/llm_translation/endpoints_client.py b/tests/e2e/llm_translation/endpoints_client.py index 4d2c73e7078..eb7bae2220c 100644 --- a/tests/e2e/llm_translation/endpoints_client.py +++ b/tests/e2e/llm_translation/endpoints_client.py @@ -75,6 +75,7 @@ class ResponsesRequest(BaseModel): stream: bool = False tools: list[ResponsesFunctionTool] | None = None guardrails: list[str] | None = None + safety_identifier: str | None = None cache: dict[str, bool] | None = {"no-cache": True} @@ -316,6 +317,7 @@ class EndpointsClient: *, stream: bool = False, guardrails: list[str] | None = None, + safety_identifier: str | None = None, ) -> StreamingResponse: return self._send( "/v1/responses", @@ -326,6 +328,7 @@ class EndpointsClient: instructions="You are a helpful assistant", stream=stream, guardrails=guardrails, + safety_identifier=safety_identifier, ), stream=stream, ) diff --git a/tests/e2e/llm_translation/test_responses_e2e.py b/tests/e2e/llm_translation/test_responses_e2e.py index 3fcf2d1ac05..525231de917 100644 --- a/tests/e2e/llm_translation/test_responses_e2e.py +++ b/tests/e2e/llm_translation/test_responses_e2e.py @@ -8,10 +8,14 @@ litellm-regression-tests/tests/test_inference_endpoints.py. from __future__ import annotations import json -from typing import cast +import threading +from collections.abc import Mapping +from dataclasses import dataclass, field +from types import MappingProxyType +from typing import Final, cast import pytest -from e2e_config import unique_marker +from e2e_config import PROVIDER_EDGE_ADVERTISE_HOST, PROVIDER_EDGE_BIND_HOST, unique_marker from e2e_http import ( assert_client_error, require_successful_call, @@ -26,7 +30,9 @@ from endpoints_client import ( ResponsesStreamEventType, ) from lifecycle import ResourceManager -from models import LiteLLMParamsBody +from models import ChatBody, ChatMessage, LiteLLMParamsBody +from provider_edge import LiveEdge, start_provider_edge +from provider_edge_bedrock import bedrock_signer from pydantic import BaseModel, ValidationError pytestmark = pytest.mark.e2e @@ -39,6 +45,33 @@ class _OptionalResponsesBody(BaseModel): BEDROCK_CONVERSE_BACKEND = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" +BEDROCK_EDGE_REGION: Final = "us-east-1" +BEDROCK_EDGE_MOUNT: Final = f"bedrock/{BEDROCK_EDGE_REGION}" + + +class ConverseRequestBody(BaseModel): + additionalModelRequestFields: dict[str, str] | None = None + + +@dataclass(slots=True) +class ConverseRequestCapture: + """The Converse bodies the proxy actually sent upstream, as seen by a live + edge sitting between the proxy and Bedrock.""" + + _bodies: list[ConverseRequestBody] = field(default_factory=list) + _lock: threading.Lock = field(default_factory=threading.Lock) + + def observe(self, url: str, headers: Mapping[str, str], body: bytes | None) -> None: + if body is None or "/converse" not in url: + return + with self._lock: + self._bodies.append(ConverseRequestBody.model_validate_json(body)) + + @property + def bodies(self) -> tuple[ConverseRequestBody, ...]: + with self._lock: + return tuple(self._bodies) + WEATHER_TOOL = ResponsesFunctionTool( name="get_weather", @@ -295,6 +328,53 @@ class TestResponses: arguments = WeatherArguments.model_validate(raw_arguments) assert arguments.location, f"function call arguments missing location: {function_call.arguments}" + @pytest.mark.provider_edge_host + @pytest.mark.parametrize("endpoint", ["/v1/responses", "/v1/chat/completions"]) + def test_bedrock_forwards_allowed_safety_identifier_as_additional_model_request_field( + self, endpoints_client: EndpointsClient, resources: ResourceManager, endpoint: str + ) -> None: + capture: Final = ConverseRequestCapture() + edge: Final = start_provider_edge( + LiveEdge(observe_request=capture.observe, sign=bedrock_signer(BEDROCK_EDGE_REGION)), + mounts=MappingProxyType({BEDROCK_EDGE_MOUNT: f"https://bedrock-runtime.{BEDROCK_EDGE_REGION}.amazonaws.com"}), + bind_host=PROVIDER_EDGE_BIND_HOST, + advertise_host=PROVIDER_EDGE_ADVERTISE_HOST, + ) + resources.defer(edge.shutdown) + model: Final = f"e2e-responses-{unique_marker()}" + model_id: Final = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model=BEDROCK_CONVERSE_BACKEND, + api_base=edge.edge.api_base(BEDROCK_EDGE_MOUNT), + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name=BEDROCK_EDGE_REGION, + allowed_openai_params=["safety_identifier"], + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key: Final = resources.key() + safety_identifier: Final = f"end-user-{unique_marker()}" + + if endpoint == "/v1/responses": + endpoints_client.responses(key, model, "reply with one word", safety_identifier=safety_identifier) + else: + endpoints_client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content="reply with one word")], + safety_identifier=safety_identifier, + ), + ) + + forwarded: Final = tuple(body.additionalModelRequestFields for body in capture.bodies) + assert forwarded, f"{endpoint} produced no Bedrock Converse request" + assert forwarded == ({"safety_identifier": safety_identifier},) * len(forwarded), ( + f"{endpoint} did not forward safety_identifier to Bedrock Converse on every attempt: {capture.bodies}" + ) + @pytest.mark.skip(reason="stage red: product gap, /v1/responses 500s (aresponses TypeError) on missing input instead of 400") @pytest.mark.covers("llm.responses.openai.input_validation.nonstream.works") def test_missing_input_returns_error( diff --git a/tests/e2e/management/test_key_management_e2e.py b/tests/e2e/management/test_key_management_e2e.py index 353b0f7cf09..39a9e657b8c 100644 --- a/tests/e2e/management/test_key_management_e2e.py +++ b/tests/e2e/management/test_key_management_e2e.py @@ -96,8 +96,8 @@ def _spend_until_budget_blocks(client: ManagementClient, key: str) -> None: for _ in range(40): outcome = client.chat_status(key, SPEND_MODEL, f"spend {unique_marker()}") if _is_budget_block(outcome): - assert outcome.status_code == 429, ( - f"budget refusal must be 429, got {outcome.status_code}: {outcome.body[:200]}" + assert outcome.status_code == 422, ( + f"budget refusal must be 422, got {outcome.status_code}: {outcome.body[:200]}" ) return assert outcome.ok, f"paid call failed before the budget tripped ({outcome.status_code}): {outcome.body[:300]}" diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 4b202e3c663..47ef672ebec 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -298,6 +298,7 @@ class ChatBody(BaseModel): max_completion_tokens: int | None = None temperature: float | None = None user: str | None = None + safety_identifier: str | None = None metadata: ChatMetadata | None = None reasoning_effort: str | None = None thinking: ThinkingParam | None = None @@ -976,6 +977,7 @@ class LiteLLMParamsBody(BaseModel): api_base: str | None = None api_version: str | None = None realtime_protocol: str | None = None + allowed_openai_params: list[str] | None = None aws_access_key_id: str | None = None aws_secret_access_key: str | None = None aws_region_name: str | None = None diff --git a/tests/e2e/provider_edge.py b/tests/e2e/provider_edge.py index 136b00208f7..fc10dde2a77 100644 --- a/tests/e2e/provider_edge.py +++ b/tests/e2e/provider_edge.py @@ -99,6 +99,7 @@ from provider_cache import ( SIGNATURE_HEADERS, CacheEdge, MountPolicy, + RequestSigner, is_bedrock, scoped_edge_base, split_test_segment, @@ -539,6 +540,7 @@ class ReplayEdge: @dataclass(frozen=True, slots=True) class LiveEdge: observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None + sign: RequestSigner | None = None type EdgeBackend = RecordEdge | ReplayEdge | LiveEdge | CacheEdge @@ -788,14 +790,16 @@ def _handle_live( method: str, url: str, headers: Mapping[str, str], body: bytes | None, timeout: float, cache: CacheEdge | None = None, mount: str = "", test_key: str | None = None, observe_request: Callable[[str, Mapping[str, str], bytes | None], None] | None = None, + sign: RequestSigner | None = None, ) -> EdgeOutcome: forwarded: Final = { name: value for name, value in headers.items() if name.lower() not in _REQUEST_DROPPED_HEADERS } if observe_request is not None: observe_request(url, forwarded, body) + outbound: Final = forwarded if sign is None else sign(method, url, forwarded, body) head: Final = ( - forward_stream(method, url, headers=forwarded, body=body, timeout=timeout) + forward_stream(method, url, headers=outbound, body=body, timeout=timeout) if cache is None else cache.forward(mount, method, url, forwarded, body, timeout, test_key=test_key) ) match head: @@ -871,10 +875,10 @@ def handle_edge_request( method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout, backend, mount, test_key, ) - case LiveEdge(observe_request=observe_request): + case LiveEdge(observe_request=observe_request, sign=sign): return _handle_live( method, _upstream_url(upstream_base, upstream_path, split.query), headers, body, timeout, - observe_request=observe_request, + observe_request=observe_request, sign=sign, ) case RecordEdge(): return _handle_record( diff --git a/tests/e2e/pytest.ini b/tests/e2e/pytest.ini index 97acb9ec52b..c6acd449884 100644 --- a/tests/e2e/pytest.ini +++ b/tests/e2e/pytest.ini @@ -13,3 +13,4 @@ markers = cli_determinism: drives the real claude CLI for several seconds; deselected unless E2E_CLI_DETERMINISM is set redis_chaos: load test that pauses the proxy's Redis outright mid-run; needs a proxy booted from gateway/redis_chaos_ci_config.yml on the same host, and is deselected unless E2E_REDIS_CHAOS is set mcp_oauth_live: real Linear OAuth consent via a captured browser session; deselected unless E2E_MCP_OAUTH_LIVE is set + provider_edge_host: routes provider traffic through the pytest host's edge in every fixture mode, so the gateway must reach the pytest host; deselected unless E2E_PROVIDER_EDGE_HOST_REACHABLE is set diff --git a/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py b/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py index 918739863ce..8a9be1d1385 100644 --- a/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py +++ b/tests/e2e/quota_management/budgets/test_budget_enforcement_e2e.py @@ -46,10 +46,10 @@ def _assert_budget_blocks(client: BudgetClient, key: str, *, user: str = "") -> pytest.fail("budget never enforced within the call budget") -def _assert_blocked_429(client: BudgetClient, key: str) -> StreamingResponse: +def _assert_blocked_422(client: BudgetClient, key: str) -> StreamingResponse: blocked = _assert_budget_blocks(client, key) - assert blocked.status_code == 429, ( - f"budget refusal must be 429, got {blocked.status_code}: {blocked.body[:200]}" + assert blocked.status_code == 422, ( + f"budget refusal must be 422, got {blocked.status_code}: {blocked.body[:200]}" ) return blocked @@ -60,7 +60,7 @@ class TestBudgetBlocksPerLevel: key = client.generate_key(max_budget=TINY_CAP) resources.defer(lambda: client.delete_key(key)) - _assert_blocked_429(client, key) + _assert_blocked_422(client, key) @pytest.mark.covers("quota_management.budget.team.blocks_over_limit") def test_team_budget_blocks_every_team_key(self, client: BudgetClient, resources: ResourceManager) -> None: @@ -71,10 +71,10 @@ class TestBudgetBlocksPerLevel: sibling_key = client.generate_key(team_id=team_id) resources.defer(lambda: client.delete_key(sibling_key)) - _assert_blocked_429(client, spender_key) + _assert_blocked_422(client, spender_key) sibling = _chat(client, sibling_key) - assert is_budget_block(sibling) and sibling.status_code == 429, ( - f"a sibling key on the capped team must get the same 429 budget_exceeded, " + assert is_budget_block(sibling) and sibling.status_code == 422, ( + f"a sibling key on the capped team must get the same 422 budget_exceeded, " f"got {sibling.status_code}: {sibling.body[:200]}" ) @@ -99,10 +99,10 @@ class TestBudgetBlocksPerLevel: team_key = client.generate_key(team_id=team_id, user_id=user_id) resources.defer(lambda: client.delete_key(team_key)) - _assert_blocked_429(client, first_key) + _assert_blocked_422(client, first_key) second = _chat(client, second_key) - assert is_budget_block(second) and second.status_code == 429, ( - f"the second personal key of a user over budget must get the same 429 budget_exceeded, " + assert is_budget_block(second) and second.status_code == 422, ( + f"the second personal key of a user over budget must get the same 422 budget_exceeded, " f"got {second.status_code}: {second.body[:200]}" ) team_result = _chat(client, team_key) @@ -133,7 +133,7 @@ class TestBudgetBlocksPerLevel: key = client.generate_key(team_id=team_id) resources.defer(lambda: client.delete_key(key)) - blocked = _assert_blocked_429(client, key) + blocked = _assert_blocked_422(client, key) assert f"Organization={org_id}" in blocked.body, ( f"refusal must name the org as the blocker, got: {blocked.body[:200]}" ) @@ -155,7 +155,7 @@ class TestBudgetBlocksPerLevel: teammate_key = client.generate_key(team_id=team_id, user_id=teammate_id) resources.defer(lambda: client.delete_key(teammate_key)) - _assert_blocked_429(client, member_key) + _assert_blocked_422(client, member_key) require_successful_call(_chat(client, teammate_key)) @@ -176,7 +176,7 @@ class TestKeyBudgetBlocksAcrossKeyKinds: control_key = client.generate_key(user_id=user_id) resources.defer(lambda: client.delete_key(control_key)) - _assert_blocked_429(client, capped_key) + _assert_blocked_422(client, capped_key) require_successful_call(_chat(client, control_key)) @pytest.mark.covers("quota_management.budget.key.blocks_over_limit") @@ -188,7 +188,7 @@ class TestKeyBudgetBlocksAcrossKeyKinds: control_key = client.generate_key(team_id=team_id) resources.defer(lambda: client.delete_key(control_key)) - _assert_blocked_429(client, capped_key) + _assert_blocked_422(client, capped_key) require_successful_call(_chat(client, control_key)) @pytest.mark.covers("quota_management.budget.key.blocks_over_limit") @@ -205,5 +205,5 @@ class TestKeyBudgetBlocksAcrossKeyKinds: control_key = client.generate_key(team_id=team_id, user_id=member_id) resources.defer(lambda: client.delete_key(control_key)) - _assert_blocked_429(client, capped_key) + _assert_blocked_422(client, capped_key) require_successful_call(_chat(client, control_key)) diff --git a/tests/e2e/quota_management/budgets/test_multi_window_budget_e2e.py b/tests/e2e/quota_management/budgets/test_multi_window_budget_e2e.py index e1cca0c0414..e04f857545d 100644 --- a/tests/e2e/quota_management/budgets/test_multi_window_budget_e2e.py +++ b/tests/e2e/quota_management/budgets/test_multi_window_budget_e2e.py @@ -102,7 +102,7 @@ def test_long_window_blocks_after_short_window_resets(client: BudgetClient, reso # 1. drive the key to get blocked by SHORT_WINDOW, assert it's budget error blocked = _drive_to_block(client, key) - assert blocked.status_code == 429, f"budget block was not a 429: {blocked.status_code} {blocked.body[:200]}" + assert blocked.status_code == 422, f"budget block was not a 422: {blocked.status_code} {blocked.body[:200]}" # 2. check the reset times of both budget windows after we drove to being blocked blocked_reset_at = window_reset_at(client.key_budget_windows(key), SHORT_WINDOW) diff --git a/tests/e2e/quota_management/budgets/test_team_multi_window_budget_e2e.py b/tests/e2e/quota_management/budgets/test_team_multi_window_budget_e2e.py index 1db68e6afe9..7683132776b 100644 --- a/tests/e2e/quota_management/budgets/test_team_multi_window_budget_e2e.py +++ b/tests/e2e/quota_management/budgets/test_team_multi_window_budget_e2e.py @@ -101,7 +101,7 @@ def test_team_long_window_blocks_after_short_window_resets(client: BudgetClient, # 1. drive the key to being blocked, assert its blocked by budget budget_exceeded blocked = _drive_to_block(client, key) - assert blocked.status_code == 429, f"budget block was not a 429: {blocked.status_code} {blocked.body[:200]}" + assert blocked.status_code == 422, f"budget block was not a 422: {blocked.status_code} {blocked.body[:200]}" # 2. check the the teams budget windows blocked_reset_at = window_reset_at(client.team_budget_windows(team_id), SHORT_WINDOW) diff --git a/tests/integration/management/test_partial_update_sequences.py b/tests/integration/management/test_partial_update_sequences.py index d79c145a685..64d807de39d 100644 --- a/tests/integration/management/test_partial_update_sequences.py +++ b/tests/integration/management/test_partial_update_sequences.py @@ -97,7 +97,7 @@ def test_zero_false_and_empty_values_are_not_treated_as_omission(gateway: Gatewa "POST", "/v1/chat/completions", {"model": models[0], "messages": [{"role": "user", "content": "zero budget"}]}, key=key, ) - assert denied.status_code == 429, denied.text + assert denied.status_code == 422, denied.text assert denied.json()["error"]["type"] == "budget_exceeded" gateway.post("/key/update", {"key": key, "max_budget": 1, "models": [], "metadata": {}}) info: Final = object_value(gateway.get("/key/info", {"key": key})["info"]) @@ -127,7 +127,7 @@ def test_zero_false_and_empty_values_are_not_treated_as_omission(gateway: Gatewa "POST", "/v1/chat/completions", {"model": models[0], "messages": [{"role": "user", "content": "updated zero budget"}]}, key=key, ) - assert zero_after_update.status_code == 429, zero_after_update.text + assert zero_after_update.status_code == 422, zero_after_update.text assert zero_after_update.json()["error"]["type"] == "budget_exceeded" gateway.post("/key/update", {"key": key, "max_budget": None}) assert read_rows( diff --git a/tests/integration/spend/test_cache_and_quota.py b/tests/integration/spend/test_cache_and_quota.py index 840594c1a96..d32297765f6 100644 --- a/tests/integration/spend/test_cache_and_quota.py +++ b/tests/integration/spend/test_cache_and_quota.py @@ -185,7 +185,7 @@ def test_key_budget_at_boundary_blocks_provider_then_explicit_reset_restores(gat {"model": model, "messages": [{"role": "user", "content": f"over budget {uuid.uuid4().hex}"}]}, key=key, ) - assert denied.status_code == 429 and denied.json()["error"]["type"] == "budget_exceeded", denied.text + assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text assert upstream.get("/__observations").json()["requests"] == [] assert gateway.chat(model, key=control, text=f"control {uuid.uuid4().hex}")["usage"]["total_tokens"] == 40 gateway.post("/key/update", {"key": key, "spend": 0}) @@ -205,7 +205,7 @@ def test_key_budget_at_boundary_blocks_provider_then_explicit_reset_restores(gat {"model": model, "messages": [{"role": "user", "content": f"boundary again {uuid.uuid4().hex}"}]}, key=key, ) - assert denied_again.status_code == 429 and denied_again.json()["error"]["type"] == "budget_exceeded", ( + assert denied_again.status_code == 422 and denied_again.json()["error"]["type"] == "budget_exceeded", ( denied_again.text ) assert upstream.get("/__observations").json()["requests"] == [] diff --git a/tests/local_testing/whitelisted_bedrock_models.txt b/tests/local_testing/whitelisted_bedrock_models.txt index 762d655b886..7615a540b23 100644 --- a/tests/local_testing/whitelisted_bedrock_models.txt +++ b/tests/local_testing/whitelisted_bedrock_models.txt @@ -133,3 +133,9 @@ meta.llama3-2-11b-instruct-v1:0 us.meta.llama3-2-11b-instruct-v1:0 meta.llama3-2-90b-instruct-v1:0 us.meta.llama3-2-90b-instruct-v1:0 +bedrock/ap-northeast-1/qwen.qwen3-next-80b-a3b +bedrock/ap-south-1/qwen.qwen3-next-80b-a3b +bedrock/ap-southeast-2/qwen.qwen3-next-80b-a3b +bedrock/eu-west-1/qwen.qwen3-next-80b-a3b +bedrock/eu-west-2/qwen.qwen3-next-80b-a3b +bedrock/sa-east-1/qwen.qwen3-next-80b-a3b diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index 7e4598c2e58..4b698f1258d 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -1,10 +1,12 @@ import asyncio import base64 +import importlib import json import os import sys from collections.abc import AsyncIterator from pathlib import Path +from types import ModuleType from typing import Final from unittest.mock import AsyncMock, MagicMock, Mock, patch @@ -2160,3 +2162,38 @@ async def test_404_before_session_initialization_preserves_method_not_found() -> ) assert caught.value.error.code == METHOD_NOT_FOUND assert caught.value.error.message == "Not Found" + + +@pytest.mark.parametrize("missing_module", ("mcp", "httpx2", "mcp.types", "openai.types.chat")) +def test_public_mcp_import_missing_dependency(missing_module: str) -> None: + with patch.dict(sys.modules): + for name in tuple(sys.modules): + if name.startswith(("litellm.experimental_mcp_client", "mcp.", "mcp_types.")) or name == "mcp": + del sys.modules[name] + with patch.dict(sys.modules, {missing_module: None}): + with pytest.raises(ImportError) as caught: + importlib.import_module("litellm.experimental_mcp_client.client") + + if missing_module in ("mcp", "httpx2"): + assert "pip install 'litellm[mcp]'" in str(caught.value) + assert isinstance(caught.value.__cause__, ModuleNotFoundError) + assert caught.value.__cause__.name == missing_module + else: + assert isinstance(caught.value, ModuleNotFoundError) + assert caught.value.name == missing_module + assert caught.value.__cause__ is None + assert "litellm[mcp]" not in str(caught.value) + + +def test_public_mcp_import_preserves_incompatible_sdk_error() -> None: + with patch.dict(sys.modules): + for name in tuple(sys.modules): + if name.startswith("litellm.experimental_mcp_client"): + del sys.modules[name] + with patch.dict(sys.modules, {"mcp": ModuleType("mcp")}): + with pytest.raises(ImportError, match="cannot import name 'ClientSession'") as caught: + importlib.import_module("litellm.experimental_mcp_client.client") + + assert not isinstance(caught.value, ModuleNotFoundError) + assert caught.value.__cause__ is None + assert "litellm[mcp]" not in str(caught.value) diff --git a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py index 626a13c8061..325052ebda9 100644 --- a/tests/test_litellm/litellm_core_utils/test_litellm_logging.py +++ b/tests/test_litellm/litellm_core_utils/test_litellm_logging.py @@ -3033,7 +3033,7 @@ def test_get_error_information_budget_exceeded_structured_fields(): assert result["error_budget_entity_id"] == "repro-user" assert result["error_budget_limit"] == 1e-06 assert result["error_budget_spend"] == 3.4e-05 - assert result["error_code"] == "429" + assert result["error_code"] == "422" assert result["error_class"] == "BudgetExceededError" assert result["error_rate_limit_type"] == "budget" @@ -6407,7 +6407,7 @@ def test_get_error_information_keeps_traceback_for_unmapped_provider_4xx(): def test_get_error_information_skips_traceback_for_budget_rejection_with_provider(): - """A key-over-budget 429 is the proxy's own rejection even after the auth + """A key-over-budget 422 is the proxy's own rejection even after the auth handler stamps the requested model's provider onto it, so it stays cheap.""" from litellm.litellm_core_utils.litellm_logging import StandardLoggingPayloadSetup @@ -6416,7 +6416,7 @@ def test_get_error_information_skips_traceback_for_budget_rejection_with_provide litellm.BudgetExceededError(current_cost=0.01, max_budget=0.0, llm_provider="anthropic") ) result = StandardLoggingPayloadSetup.get_error_information(over_budget) - assert result["error_code"] == "429" + assert result["error_code"] == "422" assert result["llm_provider"] == "anthropic" assert result["traceback"] == "" diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py index a944afc6152..6246f502344 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/adapters/test_handler_output_config_passthrough.py @@ -110,6 +110,19 @@ class TestOutputConfigStrippedFromCompletionKwargs: "reject it with 400 'Extra inputs are not permitted'" ) + def test_safeguards_is_stripped_for_non_anthropic_target(self): + extra_kwargs = { + "custom_llm_provider": "azure", + "safeguards": [{"type": "dangerous_tool_use", "classifier_context": {"v": 1}}], + } + + result = _call_prepare(extra_kwargs=extra_kwargs) + + completion_kwargs = result[0] if isinstance(result, tuple) else result + assert "safeguards" not in completion_kwargs, ( + "safeguards is an Anthropic-only field; OpenAI-format backends reject it with 400" + ) + def test_output_config_format_translated_to_response_format(self): """When ``output_config`` carries structured-output ``format``, the translator now maps it to OpenAI's ``response_format`` so non-Anthropic diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py index 997a97c6fd3..e8bfcb86bf6 100644 --- a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_anthropic_experimental_pass_through_messages_handler.py @@ -1438,3 +1438,109 @@ async def test_anthropic_messages_leaves_non_provider_failures_unmapped(): ) assert "Traceback" not in str(excinfo.value) + + +@pytest.mark.asyncio +async def test_anthropic_messages_forwards_safeguards_and_unknown_beta_to_anthropic(): + """Shapes are what Claude Code 2.1.278 sends and api.anthropic.com returns, captured 2026-09-21.""" + from litellm.llms.anthropic.experimental_pass_through.messages import handler + + safeguards = [{"type": "dangerous_tool_use", "classifier_context": {"v": 1, "permission_mode": "auto"}}] + client_betas = "dangerous-tool-use-2026-09-03,interleaved-thinking-2025-05-14" + safeguard_results = [{"type": "dangerous_tool_use", "status": {"type": "available", "tool_uses": {}}}] + captured: dict[str, object] = {} + + def upstream_records_the_request(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + captured["anthropic-beta"] = request.headers.get("anthropic-beta") + return httpx.Response( + 200, + json={ + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "claude-haiku-4-5", + "content": [{"type": "text", "text": "ok"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 1}, + "safeguard_results": safeguard_results, + }, + request=request, + ) + + upstream = AsyncHTTPHandler() + upstream.client = httpx.AsyncClient(transport=httpx.MockTransport(upstream_records_the_request)) + + response = await handler.anthropic_messages( + max_tokens=16, + messages=[{"role": "user", "content": "hi"}], + model="anthropic/claude-haiku-4-5", + custom_llm_provider="anthropic", + api_key="sk-test", + client=upstream, + safeguards=safeguards, + extra_headers={"anthropic-beta": client_betas}, + ) + + assert captured["body"]["safeguards"] == safeguards + assert set(captured["anthropic-beta"].split(",")) == set(client_betas.split(",")) + assert response["safeguard_results"] == safeguard_results + + +@pytest.mark.asyncio +async def test_anthropic_messages_streaming_forwards_safeguards_and_keeps_safeguard_results(): + """Shapes are what Claude Code 2.1.278 sends and api.anthropic.com returns, captured 2026-09-21.""" + from litellm.llms.anthropic.experimental_pass_through.messages import handler + + safeguards = [{"type": "dangerous_tool_use", "classifier_context": {"v": 1, "permission_mode": "auto"}}] + tool_verdicts = {"toolu_01": {"type": "evaluated", "outcome": "not_flagged"}} + safeguard_results = [{"type": "dangerous_tool_use", "status": {"type": "available", "tool_uses": tool_verdicts}}] + captured: dict[str, object] = {} + message_start = { + "type": "message_start", + "message": { + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "claude-haiku-4-5", + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 0}, + "safeguard_results": safeguard_results, + }, + } + message_delta = { + "type": "message_delta", + "delta": {"stop_reason": "end_turn", "stop_sequence": None, "safeguard_results": safeguard_results}, + "usage": {"output_tokens": 1}, + } + sse = "".join( + f"event: {event['type']}\ndata: {json.dumps(event)}\n\n" + for event in (message_start, message_delta, {"type": "message_stop"}) + ) + + def upstream_streams_safeguard_results(request: httpx.Request) -> httpx.Response: + captured["body"] = json.loads(request.content) + return httpx.Response(200, headers={"content-type": "text/event-stream"}, content=sse.encode(), request=request) + + upstream = AsyncHTTPHandler() + upstream.client = httpx.AsyncClient(transport=httpx.MockTransport(upstream_streams_safeguard_results)) + + stream = await handler.anthropic_messages( + max_tokens=16, + messages=[{"role": "user", "content": "hi"}], + model="anthropic/claude-haiku-4-5", + custom_llm_provider="anthropic", + api_key="sk-test", + client=upstream, + stream=True, + safeguards=safeguards, + ) + raw = b"".join([chunk async for chunk in stream]).decode() + events = [json.loads(line[len("data: ") :]) for line in raw.splitlines() if line.startswith("data: ")] + + assert captured["body"]["safeguards"] == safeguards + assert events[0]["message"]["safeguard_results"] == safeguard_results + assert [e for e in events if e["type"] == "message_delta"][0]["delta"]["safeguard_results"] == safeguard_results diff --git a/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py b/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py index 55656b97c57..27e78d35c69 100644 --- a/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py +++ b/tests/test_litellm/llms/azure_ai/image_generation/test_mai_image_generation.py @@ -453,6 +453,38 @@ class TestAzureMAIImageGeneration: ) assert round(cost, 10) == round(expected_cost, 10) + def test_mai_image_pro_edit_cost_splits_text_and_image_input(self, monkeypatch): + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + litellm.model_cost = litellm.get_model_cost_map(url="") + model = "azure_ai/MAI-Image-2.5-Pro" + model_info = litellm.get_model_info(model=model, custom_llm_provider="azure_ai") + text_tokens = 37 + image_tokens = 1024 + output_image_tokens = 1024 + + image_response = ImageResponse( + data=[ImageObject(b64_json="img1")], + usage=ImageUsage( + input_tokens=text_tokens + image_tokens, + input_tokens_details=ImageUsageInputTokensDetails( + text_tokens=text_tokens, + image_tokens=image_tokens, + ), + output_tokens=output_image_tokens, + total_tokens=text_tokens + image_tokens + output_image_tokens, + ), + ) + + cost = azure_ai_image_cost_calculator(model=model, image_response=image_response) + + expected_cost = ( + text_tokens * model_info["input_cost_per_token"] + + image_tokens * model_info["input_cost_per_image_token"] + + output_image_tokens * model_info["output_cost_per_image_token"] + ) + assert round(cost, 10) == round(expected_cost, 10) + assert model_info["input_cost_per_image_token"] != model_info["input_cost_per_token"] + def test_mai_image_cost_calculator_falls_back_to_flat_image_pricing(self, monkeypatch): monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") litellm.model_cost = litellm.get_model_cost_map(url="") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py index 4380df194ed..087c5a03498 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/auth/test_user_api_key_auth_mcp.py @@ -6339,15 +6339,15 @@ class TestMCPDcrBridgeDelegateAdmission: ) return exc_info.value - async def test_over_budget_admission_surfaces_429_not_401(self): - """A validly-authenticated but over-budget identity surfaces the standard pipeline's 429, not + async def test_over_budget_admission_surfaces_422_not_401(self): + """A validly-authenticated but over-budget identity surfaces the standard pipeline's 422, not a misleading 401. Flattening budget to 401 told the caller their credential was invalid, which on a DCR client reads as broken auth and triggers a re-authorize that cannot fix a budget problem. Regression for the status-flattening finding on the live-policy gate.""" import litellm mapped = await self._enforce_with_gate_error(litellm.BudgetExceededError(current_cost=10.0, max_budget=1.0)) - assert mapped.status_code == 429 + assert mapped.status_code == 422 async def test_db_outage_during_policy_surfaces_503_not_401(self): """A transient database outage during the live-policy gate surfaces a retryable 503, not a 401 diff --git a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py index 125b8862dfc..3edc57af124 100644 --- a/tests/test_litellm/proxy/auth/test_auth_exception_handler.py +++ b/tests/test_litellm/proxy/auth/test_auth_exception_handler.py @@ -448,7 +448,7 @@ async def test_handle_authentication_error_budget_exceeded(): ) assert exc_info.value.type == ProxyErrorTypes.budget_exceeded - assert int(exc_info.value.code) == status.HTTP_429_TOO_MANY_REQUESTS + assert int(exc_info.value.code) == status.HTTP_422_UNPROCESSABLE_CONTENT @pytest.mark.asyncio @@ -687,7 +687,7 @@ def _http_request(client_host: str | None = "10.1.2.3", headers: dict[str, str] {"allow_requests_on_db_unavailable": False}, {}, "10.1.2.3", - id="429_budget_exceeded", + id="422_budget_exceeded", ), ], ) @@ -697,7 +697,7 @@ async def test_auth_failure_logs_requester_ip_address( request_kwargs: dict[str, dict[str, str]], expected_ip: str, ) -> None: - """401s and budget 429s are rejected before `add_litellm_data_to_request` stamps + """401s and budget 422s are rejected before `add_litellm_data_to_request` stamps the caller IP, so without this the failure logs (spend logs, prometheus client_ip) had no IP, and a 401 rarely carries a key or user identity either.""" with ( diff --git a/tests/test_litellm/proxy/auth/test_handle_jwt.py b/tests/test_litellm/proxy/auth/test_handle_jwt.py index 15defb196af..e768139f04a 100644 --- a/tests/test_litellm/proxy/auth/test_handle_jwt.py +++ b/tests/test_litellm/proxy/auth/test_handle_jwt.py @@ -7144,3 +7144,52 @@ async def test_admin_jwt_team_header_only_provisions_during_admission(monkeypatc else: create_team.assert_not_awaited() assert result["team_id"] is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("existing_user", [False, True]) +@pytest.mark.parametrize("warm_cache", [False, True]) +@pytest.mark.parametrize("email", [None, "admin@external.example", "admin@allowed.example"]) +async def test_scope_admin_admission_resolves_existing_user_without_provisioning( + monkeypatch: pytest.MonkeyPatch, existing_user: bool, warm_cache: bool, email: str | None +) -> None: + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + private_key, jwk = _get_rsa_key_and_jwk("admin-status") + cache: Final = UserApiKeyCache() + cache.set_cache("litellm_jwt_auth_keys_https://admin.example/jwks", [jwk]) + user_id: Final = f"admin-status-{existing_user}-{warm_cache}-{email}" + user: Final = LiteLLM_UserTable(user_id=user_id, user_email="admin@allowed.example", metadata={"scim_active": False}, organization_memberships=[]) + if existing_user and warm_cache: + cache.set_cache(user_id, user) + database: Final = MagicMock() + users: Final = database.db.litellm_usertable + users.find_unique = AsyncMock(return_value=user if existing_user else None) + users.find_first = AsyncMock(return_value=None) + users.create = AsyncMock() + handler: Final = JWTHandler() + handler.update_environment( + prisma_client=database, + user_api_key_cache=cache, + litellm_jwtauth=LiteLLM_JWTAuth( + user_id_jwt_field="sub", user_id_upsert=True, user_email_jwt_field="email", + user_allowed_email_domain="allowed.example", + ), + ) + monkeypatch.setenv("JWT_PUBLIC_KEY_URL", "https://admin.example/jwks") + monkeypatch.setenv("JWT_ISSUER", "https://admin.example") + monkeypatch.setenv("JWT_AUDIENCE", "gateway") + token: Final = _encode_rsa_jwt( + private_key, "https://admin.example", "gateway", "admin-status", + {"sub": user_id, "scope": "litellm_proxy_admin", **({"email": email} if email else {})}, + ) + result: Final = await JWTAuthManager.auth_builder( + api_key=token, jwt_handler=handler, prisma_client=database, user_api_key_cache=cache, + parent_otel_span=None, proxy_logging_obj=MagicMock(), request_data={}, general_settings={}, route="/user/info", + ) + assert result["is_proxy_admin"] is True + assert result["user_id"] == user_id + assert result["user_object"] == (user if existing_user else None) + users.create.assert_not_awaited() + if existing_user: + assert users.find_unique.await_count == (0 if warm_cache else 1) diff --git a/tests/test_litellm/proxy/auth/test_multi_budget_windows.py b/tests/test_litellm/proxy/auth/test_multi_budget_windows.py index 0f01391b2f5..1c928448bd8 100644 --- a/tests/test_litellm/proxy/auth/test_multi_budget_windows.py +++ b/tests/test_litellm/proxy/auth/test_multi_budget_windows.py @@ -75,7 +75,7 @@ async def test_over_first_window_raises(): await _virtual_key_multi_budget_check(valid_token=token) err = exc_info.value - assert err.status_code == 429 + assert err.status_code == 422 assert "24h" in str(err) assert "Key over" in str(err) @@ -107,7 +107,7 @@ async def test_over_second_window_raises(): await _virtual_key_multi_budget_check(valid_token=token) err = exc_info.value - assert err.status_code == 429 + assert err.status_code == 422 assert "30d" in str(err) diff --git a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py index 8593be751fa..da36071a5b4 100644 --- a/tests/test_litellm/proxy/auth/test_user_api_key_auth.py +++ b/tests/test_litellm/proxy/auth/test_user_api_key_auth.py @@ -2091,7 +2091,8 @@ async def test_auto_register_binds_api_key_to_token_hash(): @pytest.mark.asyncio -async def test_auto_register_first_request_propagates_user_email(): +@pytest.mark.parametrize("active", [True, False]) +async def test_auto_register_first_request_propagates_user_email(active: bool) -> None: """ The first auto-registered JWT request must also carry user_email (resolved from the validated LiteLLM_UserTable), so attribution is consistent with the @@ -2120,6 +2121,7 @@ async def test_auto_register_first_request_propagates_user_email(): user_id="validated-user", user_email="validated@example.com", user_role="internal_user", + metadata={"scim_active": active}, ) mock_jwt_result = { "is_proxy_admin": False, @@ -2150,7 +2152,7 @@ async def test_auto_register_first_request_propagates_user_email(): patch("litellm.proxy.proxy_server.master_key", "sk-master"), patch("litellm.proxy.proxy_server.prisma_client", prisma_client), patch("litellm.proxy.proxy_server.user_api_key_cache", user_api_key_cache), - patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock(post_call_failure_hook=AsyncMock(return_value=None))), patch("litellm.proxy.proxy_server.jwt_handler", jwt_handler), patch( "litellm.proxy.auth.user_api_key_auth._resolve_jwt_to_virtual_key", @@ -2170,8 +2172,22 @@ async def test_auto_register_first_request_propagates_user_email(): "litellm.proxy.auth.user_api_key_auth._auto_register_jwt_mapping", new_callable=AsyncMock, return_value=auto_registered_key, - ), + ) as auto_register, ): + if not active: + with pytest.raises(ProxyException, match="deactivated via SCIM") as exc: + await _user_api_key_auth_builder( + request=mock_request, + api_key=jwt_token, + azure_api_key_header="", + anthropic_api_key_header=None, + google_ai_studio_api_key_header=None, + azure_apim_header=None, + request_data={}, + ) + assert int(exc.value.code) == 401 + auto_register.assert_not_awaited() + return result = await _user_api_key_auth_builder( request=mock_request, api_key=jwt_token, @@ -7315,15 +7331,15 @@ class TestJWTAuthUserEmail: the Prometheus `user_email` label and `user_api_key_user_email` in StandardLogging/SpendLogs metadata, which were always None for JWT traffic.""" - def _jwt_request(self, jwt_token): + def _jwt_request(self, jwt_token, route="/v1/chat/completions"): mock_request = MagicMock() - mock_request.url.path = "/v1/chat/completions" - mock_request.method = "POST" + mock_request.url.path = route + mock_request.method = "GET" if route.endswith("/list") else "POST" mock_request.headers = {"authorization": f"Bearer {jwt_token}"} mock_request.query_params = {} return mock_request - async def _run_jwt_auth(self, mock_jwt_result, jwt_token): + async def _run_jwt_auth(self, mock_jwt_result, jwt_token, route="/v1/chat/completions"): with ( patch( "litellm.proxy.proxy_server.general_settings", @@ -7344,7 +7360,7 @@ class TestJWTAuthUserEmail: litellm_jwtauth=LiteLLM_JWTAuth(), ) return await user_api_key_auth( - request=self._jwt_request(jwt_token), + request=self._jwt_request(jwt_token, route), api_key=f"Bearer {jwt_token}", ) @@ -7376,6 +7392,44 @@ class TestJWTAuthUserEmail: assert result.user_id == "jwt-human-user" assert result.user_email == "resolved@example.com" + @pytest.mark.asyncio + @pytest.mark.parametrize("route", ["/mcp-rest/tools/list", "/mcp-rest/tools/call", "/v1/chat/completions", "/user/info"]) + @pytest.mark.parametrize("active", [False, True, None, "false", 0]) + @pytest.mark.parametrize("is_admin", [False, True]) + async def test_jwt_auth_rejects_deactivated_user( + self, route: str, active: bool | str | int | None, is_admin: bool + ) -> None: + from typing import Final + + jwt_token: Final = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" + result: Final = { + "is_proxy_admin": is_admin, + "team_object": None, + "user_object": LiteLLM_UserTable( + user_id="jwt-human-user", + user_role=LitellmUserRoles.PROXY_ADMIN.value if is_admin else LitellmUserRoles.INTERNAL_USER.value, + metadata={} if active is None else {"scim_active": active}, + ), + "end_user_object": None, + "org_object": None, + "token": jwt_token, + "team_id": None, + "user_id": "jwt-human-user", + "user_email": None, + "end_user_id": None, + "org_id": None, + "team_membership": None, + "jwt_claims": {"sub": "user1"}, + } + + if active is False: + with pytest.raises(ProxyException, match="deactivated via SCIM") as exc: + await self._run_jwt_auth(result, jwt_token, route) + assert int(exc.value.code) == 401 + else: + token: Final = await self._run_jwt_auth(result, jwt_token, route) + assert token.user_id == "jwt-human-user" + @pytest.mark.asyncio async def test_jwt_auth_populates_user_email_on_proxy_admin(self): jwt_token = "eyJhbGciOiJSUzI1NiJ9.eyJzdWIiOiJ1c2VyMSJ9.signature" diff --git a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py index 0a9cf8b84cb..2cfbfbbb3cb 100644 --- a/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py +++ b/tests/test_litellm/proxy/management_endpoints/scim/test_scim_key_deactivation.py @@ -541,3 +541,63 @@ async def test_scim_put_user_explicit_active_false_blocks_keys(): assert update_kwargs["where"] == {"token": "hash-block-me"} assert update_kwargs["data"]["blocked"] is True assert '"scim_blocked": true' in update_kwargs["data"]["metadata"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize("method", ["PUT", "PATCH"]) +@pytest.mark.parametrize("active", [False, True]) +@pytest.mark.parametrize("failure", [None, "write", "keys"]) +@pytest.mark.parametrize("status_change", [False, True]) +async def test_scim_status_write_refreshes_user_cache( + method: str, active: bool, failure: str | None, status_change: bool +) -> None: + import json + from typing import Final + + from litellm.proxy._types import ProxyException + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + user_id: Final = "scim-cache-user" + saved: Final = LiteLLM_UserTable( + user_id=user_id, user_email="x@example.com", teams=[], metadata={"scim_active": not active if status_change else active}, + ) + updated: Final = LiteLLM_UserTable( + user_id=user_id, user_email="x@example.com", teams=[], metadata={"scim_active": active}, + ) + client, db = _build_prisma_with_keys([], mock_user=saved.model_copy(deep=True), updated_user=updated) + if failure == "write": + db.litellm_usertable.update.side_effect = RuntimeError("status write failed") + if failure == "keys": + db.litellm_verificationtoken.find_many.side_effect = RuntimeError("key update failed") + cache: Final = UserApiKeyCache() + await cache.async_set_cache(key=user_id, value=saved, model_type=LiteLLM_UserTable) + with ( + patch("litellm.proxy.proxy_server.prisma_client", client), # test-quality-ok: substitute the database dependency + patch("litellm.proxy.proxy_server.user_api_key_cache", cache), # test-quality-ok: exercise a real isolated cache + patch("litellm.proxy.proxy_server.proxy_logging_obj", MagicMock()), # test-quality-ok: isolate the logging dependency + patch( # test-quality-ok: observe the Redis publication boundary + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation", + new_callable=AsyncMock, + ) as broadcast, + ): + request: Final = ( + update_user(user_id=user_id, user=SCIMUser.model_validate(_build_put_user_payload(user_id, active=active))) + if method == "PUT" else + patch_user(user_id=user_id, patch_ops=SCIMPatchOp( + Operations=[SCIMPatchOperation(op="replace", path="active", value=active)] + )) + ) + if failure == "write" or (failure == "keys" and status_change): + with pytest.raises(ProxyException, match="status write failed" if failure == "write" else "key update failed"): + await request + else: + response: Final = await request + assert response.active is active + assert json.loads(db.litellm_usertable.update.await_args.kwargs["data"]["metadata"])["scim_active"] is active + cached: Final = await cache.async_get_cache(key=user_id, model_type=LiteLLM_UserTable) + if failure == "write": + assert cached == saved + broadcast.assert_not_awaited() + else: + assert cached is None + broadcast.assert_awaited_once_with(cache_key=user_id) diff --git a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py index 3f2ba365a04..b64f7c8fb7c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_internal_user_endpoints.py @@ -2228,6 +2228,51 @@ async def test_user_model_budget_update_by_email_refreshes_cached_user(mocker: M broadcast.assert_awaited_once_with(cache_key=saved_user.user_id) +@pytest.mark.asyncio +@pytest.mark.parametrize("by_email", [False, True]) +@pytest.mark.parametrize("active", [False, True, None]) +async def test_user_status_update_refreshes_cached_user( + mocker: MockerFixture, by_email: bool, active: bool | None +) -> None: + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.management_endpoints.internal_user_endpoints import _update_single_user_helper + + saved_user: Final = LiteLLM_UserTable( + user_id="user-spruce", + user_email="spruce@example.test", + metadata={"scim_active": False if active is None else not active, "department": "engineering"}, + ) + prisma_client: Final = mocker.MagicMock() + prisma_client.db.litellm_usertable.find_first = mocker.AsyncMock(return_value=saved_user) + prisma_client.get_data = mocker.AsyncMock(return_value=[saved_user]) + prisma_client.update_data = mocker.AsyncMock(return_value={"user_id": saved_user.user_id, "data": saved_user}) + mocker.patch("litellm.proxy.proxy_server.prisma_client", prisma_client) # test-quality-ok: substitute the database dependency + cache: Final = UserApiKeyCache() + await cache.async_set_cache(key=saved_user.user_id, value=saved_user, model_type=LiteLLM_UserTable) + mocker.patch("litellm.proxy.proxy_server.user_api_key_cache", cache) # test-quality-ok: exercise a real isolated cache + broadcast: Final = mocker.patch( # test-quality-ok: observe the Redis publication boundary + "litellm.proxy.common_utils.auth_cache_invalidation_pubsub.publish_auth_cache_invalidation", + new_callable=mocker.AsyncMock, + ) + + await _update_single_user_helper( + user_request=UpdateUserRequest( + user_id=None if by_email else saved_user.user_id, + user_email=saved_user.user_email if by_email else None, + metadata={"department": "engineering"} if active is None else {"scim_active": active}, + ), + user_api_key_dict=UserAPIKeyAuth(user_id="admin-spruce", user_role=LitellmUserRoles.PROXY_ADMIN), + ) + + assert prisma_client.update_data.call_args.kwargs["user_id"] == saved_user.user_id + assert prisma_client.update_data.call_args.kwargs["data"]["metadata"] == ( + {"department": "engineering"} if active is None else {"scim_active": active} + ) + assert await cache.async_get_cache(key=saved_user.user_id, model_type=LiteLLM_UserTable) is None + broadcast.assert_awaited_once_with(cache_key=saved_user.user_id) + + @pytest.mark.asyncio async def test_bulk_user_model_budget_clear_serializes_and_refreshes_cache(mocker: MockerFixture) -> None: from litellm.proxy._types import LiteLLM_UserTable diff --git a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py index e2a68988ee2..b02ff47ed52 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_key_management_endpoints.py @@ -8156,7 +8156,7 @@ async def test_reset_key_spend_resets_budget_windows(monkeypatch): counter without also advancing reset_at is not durable either: the very next request would re-sum the unchanged historical spend and put the counter right back above the window's max_budget, so - _virtual_key_multi_budget_check kept raising BudgetExceededError (429) on + _virtual_key_multi_budget_check kept raising BudgetExceededError (422) on every request even though the key's own reported spend read $0. """ mock_prisma_client = MagicMock() @@ -16593,7 +16593,7 @@ async def test_info_key_fn_reads_the_configured_budget_model_key(monkeypatch): It used to probe a second, provider-stripped key because the counter was written under the request model instead, which is what let a key report zero - usage while being blocked at 429. + usage while being blocked at 422. """ from unittest.mock import AsyncMock, MagicMock diff --git a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py index f7abb209015..4153bf7d7ee 100644 --- a/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py +++ b/tests/test_litellm/proxy/response_api_endpoints/test_endpoints.py @@ -2099,7 +2099,7 @@ class TestCursorVariantPerModelBudgetEnforcement: response = _post_cursor_with_real_auth(valid_token, attrs, request_model="claude-opus-5-thinking-high") - assert response.status_code == 429, response.text + assert response.status_code == 422, response.text error = response.json()["error"] assert error["type"] == "budget_exceeded" assert "exceeded budget for model=claude-opus-5" in error["message"] @@ -2110,8 +2110,8 @@ class TestCursorVariantPerModelBudgetEnforcement: base_response = _post_cursor_with_real_auth(valid_token, attrs, request_model="claude-opus-5") alias_response = _post_cursor_with_real_auth(valid_token, attrs, request_model="claude-opus-5-fast") - assert base_response.status_code == 429, base_response.text - assert alias_response.status_code == 429, alias_response.text + assert base_response.status_code == 422, base_response.text + assert alias_response.status_code == 422, alias_response.text assert alias_response.json() == base_response.json() diff --git a/tests/test_litellm/proxy/test_common_request_processing.py b/tests/test_litellm/proxy/test_common_request_processing.py index e4ca0b03d59..0b872400be0 100644 --- a/tests/test_litellm/proxy/test_common_request_processing.py +++ b/tests/test_litellm/proxy/test_common_request_processing.py @@ -495,7 +495,7 @@ class TestProxyBaseLLMRequestProcessing: ) assert exc_info.value.type == ProxyErrorTypes.budget_exceeded - assert exc_info.value.code == "429" + assert exc_info.value.code == "422" tag_budget_check.assert_awaited_once() _, call_kwargs = tag_budget_check.call_args assert call_kwargs["tags"] == ("guardrail-tag",) @@ -702,7 +702,7 @@ class TestProxyBaseLLMRequestProcessing: ) assert exc_info.value.type == ProxyErrorTypes.budget_exceeded - assert exc_info.value.code == "429" + assert exc_info.value.code == "422" assert "guardrail-tag" in exc_info.value.message @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index f71f9c20f3b..950a6cc3c40 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -10872,7 +10872,7 @@ async def test_realtime_session_rejected_in_pre_call_releases_the_budget_reserva """A rate-limit or guardrail rejection happens before route_request, so the relay never runs and no success log can own the reservation. The endpoint must release it on that exit too, or the key stays pinned at the reserved - amount and its next requests 429 with budget_exceeded while /key/info shows + amount and its next requests 422 with budget_exceeded while /key/info shows spend 0 (reproduced live with rpm_limit=1). The client still gets the pre-call error event and the 1011 close it got before.""" reservation: Final = {"reserved_cost": 0.55, "input_cost": 0.0, "finalized": False, "entries": []} diff --git a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py index 6077f281a81..7b9de4644b4 100644 --- a/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py +++ b/tests/test_litellm/responses/litellm_completion_transformation/test_litellm_completion_responses.py @@ -1248,6 +1248,16 @@ class TestFunctionCallTransformation: assert "tool_choice" not in result assert "tools" not in result + def test_safety_identifier_forwarded_to_chat_completion_request(self) -> None: + result: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request( + model="bedrock/global.openai.gpt-5.6-luna", + input="hi", + responses_api_request={"safety_identifier": "user-7f3a"}, + custom_llm_provider="bedrock", + ) + + assert result["safety_identifier"] == "user-7f3a" + def test_parallel_tool_calls_dropped_when_no_chat_tools_remain(self) -> None: transform: Final = LiteLLMCompletionResponsesConfig.transform_responses_api_request_to_chat_completion_request codex_tool_search: Final = { diff --git a/tests/test_litellm/test_cost_calculator.py b/tests/test_litellm/test_cost_calculator.py index aef17f3d5d0..da3d022d669 100644 --- a/tests/test_litellm/test_cost_calculator.py +++ b/tests/test_litellm/test_cost_calculator.py @@ -4230,3 +4230,30 @@ def test_completion_cost_prices_responses_websocket_turns_per_service_tier(): assert ws_cost == pytest.approx(_http_cost(100, 40, "default") + _http_cost(60, 10, "priority")) assert ws_cost != pytest.approx(_http_cost(160, 50, "default")) assert ws_cost != pytest.approx(_http_cost(160, 50, "priority")) + + +QWEN3_NEXT_REGIONS: Final = ("ap-northeast-1", "ap-south-1", "ap-southeast-2", "eu-west-1", "eu-west-2", "sa-east-1") + + +@pytest.mark.parametrize("region", QWEN3_NEXT_REGIONS) +def test_cost_per_token_bedrock_qwen3_next_uses_regional_entry_not_us_rate( + monkeypatch: pytest.MonkeyPatch, region: str +) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + + regional: Final = litellm.model_cost[f"bedrock/{region}/qwen.qwen3-next-80b-a3b"] + us: Final = litellm.model_cost["qwen.qwen3-next-80b-a3b"] + assert regional["input_cost_per_token"] != us["input_cost_per_token"] + assert regional["output_cost_per_token"] != us["output_cost_per_token"] + + prompt_tokens, completion_tokens = 1000, 500 + prompt_usd, completion_usd = cost_per_token( + model=f"bedrock/{region}/qwen.qwen3-next-80b-a3b", + prompt_tokens=prompt_tokens, + completion_tokens=completion_tokens, + custom_llm_provider="bedrock", + ) + + assert prompt_usd == pytest.approx(prompt_tokens * regional["input_cost_per_token"]) + assert completion_usd == pytest.approx(completion_tokens * regional["output_cost_per_token"]) diff --git a/tests/test_litellm/test_rate_limit_error_unification.py b/tests/test_litellm/test_rate_limit_error_unification.py index 8241b29aff1..e5acba938c7 100644 --- a/tests/test_litellm/test_rate_limit_error_unification.py +++ b/tests/test_litellm/test_rate_limit_error_unification.py @@ -1397,13 +1397,18 @@ class TestBudgetExceededErrorSurfacesUnifiedFields: assert e.llm_provider == "anthropic" def test_should_keep_existing_status_code_and_message(self): - # Backward-compat guard: existing callers depend on `status_code=429` + # Backward-compat guard: existing callers depend on `status_code=422` # and the canonical message format. e = litellm.BudgetExceededError(current_cost=0.000109, max_budget=0.0001) - assert e.status_code == 429 + assert e.status_code == 422 assert "Current cost: 0.000109" in e.message assert "Max budget: 0.0001" in e.message + def test_should_honor_budget_exceeded_status_code_override(self, monkeypatch: pytest.MonkeyPatch): + monkeypatch.setattr(litellm, "budget_exceeded_status_code", 429) + e = litellm.BudgetExceededError(current_cost=0.5, max_budget=0.1) + assert e.status_code == 429 + def test_should_still_be_catchable_as_exception_not_rate_limit_error(self): # Critical: we deliberately did NOT make BudgetExceededError a # RateLimitError subclass. Existing `except BudgetExceededError:` @@ -1424,7 +1429,7 @@ class TestBudgetExceededErrorSurfacesUnifiedFields: info = StandardLoggingPayloadSetup.get_error_information(e) assert info["error_rate_limit_category"] == "litellm_rate_limit" assert info["error_rate_limit_type"] == "budget" - assert info["error_code"] == "429" + assert info["error_code"] == "422" assert info["error_class"] == "BudgetExceededError" def test_should_propagate_llm_provider_to_standard_logging_payload(self): diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 7a7d5d44e95..8bdda0490c0 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -652,6 +652,7 @@ def validate_model_cost_values(model_data, exceptions=None): "cache_creation_input_audio_token_cost", "cache_read_input_token_cost", "cache_read_input_audio_token_cost", + "cache_read_input_image_token_cost", "input_dbu_cost_per_token", "output_db_cost_per_token", "output_dbu_cost_per_token", @@ -740,6 +741,7 @@ def test_aaamodel_prices_and_context_window_json_is_valid(): "cache_read_input_token_cost_above_512k_tokens": {"type": "number"}, "cache_creation_input_token_cost_above_1hr_above_200k_tokens": {"type": "number"}, "cache_read_input_audio_token_cost": {"type": "number"}, + "cache_read_input_image_token_cost": {"type": "number"}, "audio_transcription_config": {"type": "string"}, "deprecation_date": {"type": "string"}, "input_cost_per_audio_per_second": {"type": "number"}, diff --git a/tests/unit/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py b/tests/unit/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py index dcea462d7a1..cc314351fc2 100644 --- a/tests/unit/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py +++ b/tests/unit/a2a_protocol/providers/bedrock_agentcore/test_bedrock_agentcore_a2a.py @@ -530,6 +530,38 @@ class TestNonStreaming: assert sent_headers.get("x-mcp-token") == "mcp-abc" +class TestStreaming: + """Streaming requests must ask AgentCore for a stream, not a single send.""" + + @pytest.mark.asyncio + async def test_streaming_request_uses_message_stream_method_and_yields_sse_events(self, httpx_transport): + from litellm.a2a_protocol.providers.bedrock_agentcore.config import ( + BedrockAgentCoreA2AConfig, + ) + + sse_body = ( + 'data: {"jsonrpc": "2.0", "id": "req-001", "result": {"kind": "task", "id": "t1"}}\n\n' + 'data: {"jsonrpc": "2.0", "id": "req-001", "result": {"kind": "status-update", "final": true}}\n\n' + ) + with respx.mock(assert_all_called=True) as router: + route = router.post(url__regex=r".*/invocations.*").mock( + return_value=httpx.Response(200, headers={"content-type": "text/event-stream"}, text=sse_body) + ) + events = [ + event + async for event in BedrockAgentCoreA2AConfig().handle_streaming( + request_id="req-001", + params=SAMPLE_PARAMS, + litellm_params=SAMPLE_LITELLM_PARAMS, + ) + ] + + sent_body = json.loads(route.calls.last.request.content) + assert sent_body["method"] == "message/stream", sent_body + assert sent_body["params"]["message"]["messageId"] == "msg-001" + assert [event["result"]["kind"] for event in events] == ["task", "status-update"] + + class TestConfigManager: """Test that config manager routes 'bedrock' correctly.""" diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index e33764c3d1a..d62a758e3a8 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -30903,6 +30903,8 @@ export interface components { cache_creation_input_token_cost_ultrafast?: number | null; /** Cache Read Input Audio Token Cost */ cache_read_input_audio_token_cost?: number | null; + /** Cache Read Input Image Token Cost */ + cache_read_input_image_token_cost?: number | null; /** Cache Read Input Token Cost */ cache_read_input_token_cost?: number | null; /** Cache Read Input Token Cost Above 200K Tokens */ @@ -41620,6 +41622,8 @@ export interface components { cache_creation_input_token_cost_ultrafast?: number | null; /** Cache Read Input Audio Token Cost */ cache_read_input_audio_token_cost?: number | null; + /** Cache Read Input Image Token Cost */ + cache_read_input_image_token_cost?: number | null; /** Cache Read Input Token Cost */ cache_read_input_token_cost?: number | null; /** Cache Read Input Token Cost Above 200K Tokens */ diff --git a/whitelisted_bedrock_models.txt b/whitelisted_bedrock_models.txt index 6124cb41044..254842e2714 100644 --- a/whitelisted_bedrock_models.txt +++ b/whitelisted_bedrock_models.txt @@ -44,6 +44,7 @@ bedrock/ap-northeast-1/minimax.minimax-m2.5 bedrock/ap-northeast-1/moonshotai.kimi-k2-thinking bedrock/ap-northeast-1/moonshotai.kimi-k2.5 bedrock/ap-northeast-1/qwen.qwen3-coder-next +bedrock/ap-northeast-1/qwen.qwen3-next-80b-a3b bedrock/moonshotai.kimi-k2-thinking bedrock/moonshotai.kimi-k2.5 bedrock/ap-south-1/meta.llama3-70b-instruct-v1:0 @@ -54,6 +55,7 @@ bedrock/ap-south-1/minimax.minimax-m2.5 bedrock/ap-south-1/moonshotai.kimi-k2-thinking bedrock/ap-south-1/moonshotai.kimi-k2.5 bedrock/ap-south-1/qwen.qwen3-coder-next +bedrock/ap-south-1/qwen.qwen3-next-80b-a3b bedrock/ap-southeast-2/minimax.minimax-m2.5 bedrock/ap-southeast-3/deepseek.v3.2 bedrock/ap-southeast-3/minimax.minimax-m2.1 @@ -83,11 +85,13 @@ bedrock/eu-west-1/meta.llama3-8b-instruct-v1:0 bedrock/eu-west-1/minimax.minimax-m2.1 bedrock/eu-west-1/minimax.minimax-m2.5 bedrock/eu-west-1/qwen.qwen3-coder-next +bedrock/eu-west-1/qwen.qwen3-next-80b-a3b bedrock/eu-west-2/meta.llama3-70b-instruct-v1:0 bedrock/eu-west-2/meta.llama3-8b-instruct-v1:0 bedrock/eu-west-2/minimax.minimax-m2.1 bedrock/eu-west-2/minimax.minimax-m2.5 bedrock/eu-west-2/qwen.qwen3-coder-next +bedrock/eu-west-2/qwen.qwen3-next-80b-a3b bedrock/eu-west-3/mistral.mistral-7b-instruct-v0:2 bedrock/eu-west-3/mistral.mistral-large-2402-v1:0 bedrock/eu-west-3/mistral.mixtral-8x7b-instruct-v0:1 @@ -103,6 +107,7 @@ bedrock/sa-east-1/minimax.minimax-m2.5 bedrock/sa-east-1/moonshotai.kimi-k2-thinking bedrock/sa-east-1/moonshotai.kimi-k2.5 bedrock/sa-east-1/qwen.qwen3-coder-next +bedrock/sa-east-1/qwen.qwen3-next-80b-a3b bedrock/us-east-1/1-month-commitment/anthropic.claude-instant-v1 bedrock/us-east-1/1-month-commitment/anthropic.claude-v1 bedrock/us-east-1/1-month-commitment/anthropic.claude-v2:1 @@ -240,3 +245,4 @@ bedrock/us-gov-east-1/anthropic.claude-sonnet-5 bedrock/us-gov-east-1/anthropic.claude-opus-4-8 bedrock/us-gov-east-1/anthropic.claude-opus-5 bedrock/us-gov-east-1/anthropic.claude-fable-5-1 +bedrock/ap-southeast-2/qwen.qwen3-next-80b-a3b