diff --git a/.github/ci-coverage-allowlist.yml b/.github/ci-coverage-allowlist.yml index 0da07038152..672f102eeb1 100644 --- a/.github/ci-coverage-allowlist.yml +++ b/.github/ci-coverage-allowlist.yml @@ -106,11 +106,6 @@ dockerfiles: and lint workflows already exercise that output, so building the image adds no signal about it paths: - ui/Dockerfile - - reason: >- - The Rust gateway ships as its own chart and package with a separate release pipeline, so its - image is not part of this repo's Python image set - paths: - - litellm-rust/crates/ai-gateway/Dockerfile - reason: >- An example image under cookbook/ that is documentation rather than a shipped artifact paths: diff --git a/.github/scripts/verify_linux_native_wheel.py b/.github/scripts/verify_linux_native_wheel.py index 899e2a211c0..4fb8f068eb0 100644 --- a/.github/scripts/verify_linux_native_wheel.py +++ b/.github/scripts/verify_linux_native_wheel.py @@ -205,7 +205,7 @@ def main( native_module: Final = load_native_module(native_path) native_module_loads: Final = native_module is not None panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test") - native_size_limit: Final = 20_000_000 + native_size_limit: Final = 25_000_000 native_size_within_limit: Final = native_member.file_size <= native_size_limit validations: Final = ( (f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG), @@ -222,7 +222,7 @@ def main( ("Python extension entry point is present", extension_entry_point_present), ("Native module loads", native_module_loads), ("Production module omits the panic test hook", panic_test_hook_absent), - ("Native extension does not exceed 20 MB", native_size_within_limit), + ("Native extension does not exceed 25 MB", native_size_within_limit), ("Wheel contents are valid", not unexpected_members), ) diff --git a/.github/workflows/ai-gateway-image.yml b/.github/workflows/ai-gateway-image.yml new file mode 100644 index 00000000000..3f690f566b0 --- /dev/null +++ b/.github/workflows/ai-gateway-image.yml @@ -0,0 +1,73 @@ +name: ai-gateway image + +on: + push: + paths: + - "litellm-rust/**" + - "litellm/**" + - "enterprise/**" + - "litellm-proxy-extras/**" + - "pyproject.toml" + - "rust-toolchain.toml" + - ".github/workflows/ai-gateway-image.yml" + pull_request: + branches: + - main + - litellm_internal_staging + - litellm_oss_staging + - "litellm_**" + paths: + - "litellm-rust/**" + - "litellm/**" + - "enterprise/**" + - "litellm-proxy-extras/**" + - "pyproject.toml" + - "rust-toolchain.toml" + - ".github/workflows/ai-gateway-image.yml" + +permissions: + contents: read + +concurrency: + group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} + cancel-in-progress: true + +jobs: + ai-gateway-image: + name: ai-gateway release image + runs-on: ubuntu-latest + timeout-minutes: 60 + permissions: + contents: read + steps: + - name: Checkout repository + uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 + with: + persist-credentials: false + - name: Build the release image + run: docker build -f litellm-rust/crates/ai-gateway/Dockerfile -t litellm-ai-gateway:${{ github.sha }} . + - name: Start the gateway and wait for readiness + env: + IMAGE: litellm-ai-gateway:${{ github.sha }} + run: | + docker run -d --name ai-gateway -p 4001:4001 \ + -e LITELLM_MASTER_KEY=sk-ci-not-a-real-key \ + -e OPENAI_API_KEY=sk-ci-not-a-real-key \ + "$IMAGE" + for _ in $(seq 1 60); do + if curl -fsS http://127.0.0.1:4001/health/readiness; then + echo "gateway is serving readiness" + exit 0 + fi + sleep 2 + done + echo "gateway never became ready" >&2 + docker logs ai-gateway >&2 + exit 1 + - name: Assert the gateway loaded the baked config + run: | + docker logs ai-gateway 2>&1 | tee gateway.log + grep 'via python config reader' gateway.log + - name: Stop the gateway + if: always() + run: docker rm -f ai-gateway || true diff --git a/README.md b/README.md index 92757fcbbc1..901cc5b0cea 100644 --- a/README.md +++ b/README.md @@ -262,6 +262,8 @@ curl -X POST 'http://0.0.0.0:4000/v1/chat/completions' \ } ``` +For MCP OAuth, an upstream may advertise dynamic client registration but refuse requests with HTTP 401 or 403. If the provider requires a pre-registered OAuth app, configure its `credentials.client_id` and, when required, `credentials.client_secret` on the MCP server. This skips dynamic registration in the gateway sign-in flow. The provider must approve the app for MCP access; reaching its authorization page does not establish that login or tool calls will succeed + [**Docs: MCP Gateway**](https://docs.litellm.ai/docs/mcp) diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20260910000000_add_autorouter_session_baseline_models/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260910000000_add_autorouter_session_baseline_models/migration.sql new file mode 100644 index 00000000000..e7ce1a3180b --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20260910000000_add_autorouter_session_baseline_models/migration.sql @@ -0,0 +1 @@ +ALTER TABLE "LiteLLM_AutoRouterSession" ADD COLUMN IF NOT EXISTS "baseline_models" JSONB NOT NULL DEFAULT '{}'; diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 817df082d8c..7d521d54791 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1514,6 +1514,7 @@ model LiteLLM_AutoRouterSession { classifier_cost Float @default(0) classifier_cost_recorded_turns Int @default(0) tier_turns Json @default("{}") + baseline_models Json @default("{}") @@id([api_key, session_id, router_name]) @@index([last_turn_at], map: "idx_autorouter_session_last_turn") diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 0e8e6e09a21..cf000e85d68 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -2,6 +2,12 @@ # It is not intended for manual editing. version = 4 +[[package]] +name = "adler2" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "320119579fcad9c21884f5c4861d16174d0e06250625266f50fe6898340abefa" + [[package]] name = "ahash" version = "0.8.12" @@ -34,6 +40,15 @@ dependencies = [ "cc", ] +[[package]] +name = "android_system_properties" +version = "0.1.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae221649c9976a6f6c56ae1facf410f3ddb33cc661c4b7b61020a912d4237fbc" +dependencies = [ + "libc", +] + [[package]] name = "anes" version = "0.1.6" @@ -55,6 +70,29 @@ dependencies = [ "rustversion", ] +[[package]] +name = "async-compression" +version = "0.4.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4f10dafd0c8d2e51ae9a748805777613ed0bbe17bf586b76c8311f45c020a32f" +dependencies = [ + "compression-codecs", + "compression-core", + "pin-project-lite", + "tokio", +] + +[[package]] +name = "async-lock" +version = "3.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "290f7f2596bd5b78a9fec8088ccd89180d7f9f55b94b0576823bbbdc72ee8311" +dependencies = [ + "event-listener", + "event-listener-strategy", + "pin-project-lite", +] + [[package]] name = "async-trait" version = "0.1.91" @@ -482,6 +520,58 @@ dependencies = [ "tracing", ] +[[package]] +name = "azure_core" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4e41cbd819986ba41904c207d8ffc4106f8f8352a548d773e9554906379bb2fb" +dependencies = [ + "async-lock", + "async-trait", + "azure_core_macros", + "bytes", + "futures", + "pin-project", + "rustc_version", + "serde", + "serde_json", + "tokio", + "tracing", + "typespec", + "typespec_client_core", +] + +[[package]] +name = "azure_core_macros" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b9b52dba6a345f3ad2d42ff8d0d63df9d0994cfa29657bf18ffdbf149f78a4f5" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", + "tracing", +] + +[[package]] +name = "azure_identity" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32edf96b356ca7c51d7590c4925cc36efc3947a5da4468e8e0b25c56ecbb3de5" +dependencies = [ + "async-lock", + "async-trait", + "azure_core", + "futures", + "pin-project", + "serde", + "serde_json", + "time", + "tokio", + "tracing", + "url", +] + [[package]] name = "base64" version = "0.13.1" @@ -494,6 +584,12 @@ version = "0.22.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" +[[package]] +name = "base64" +version = "0.23.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ac07cdecf99051d9a5238b80f35af32cdeba5b336e55d957b318b50137e18da5" + [[package]] name = "base64-simd" version = "0.8.0" @@ -606,6 +702,20 @@ dependencies = [ "rand_core 0.10.1", ] +[[package]] +name = "chrono" +version = "0.4.45" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1aa79e62e7697b8e29b513a68abacf485adcd1fe8284a4316c5ae868e6633327" +dependencies = [ + "iana-time-zone", + "js-sys", + "num-traits", + "serde", + "wasm-bindgen", + "windows-link", +] + [[package]] name = "ciborium" version = "0.2.2" @@ -673,6 +783,16 @@ version = "0.5.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0c9ea0ac24bc397ab3c98583a3c9ba74fa56b09a4449bbe172b9b1ddb016027a" +[[package]] +name = "combine" +version = "4.6.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cfc320937d09e6de266b31b9afb480f197d7a861be86be7cb2ea7e5d1bfffc5e" +dependencies = [ + "bytes", + "memchr", +] + [[package]] name = "compact_str" version = "0.9.1" @@ -688,6 +808,23 @@ dependencies = [ "static_assertions", ] +[[package]] +name = "compression-codecs" +version = "0.4.41" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "58a6d0db8759036a783bc7c3f7a07f8cef3bf9470eb1db3bc86e8bcd1c5d0fe8" +dependencies = [ + "compression-core", + "flate2", + "memchr", +] + +[[package]] +name = "compression-core" +version = "0.4.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e8ccc4ea9f6acc32d102c0f6d471d11d913ad15f20c04de743374861fa1d414" + [[package]] name = "const-oid" version = "0.10.2" @@ -728,6 +865,15 @@ dependencies = [ "libc", ] +[[package]] +name = "crc32fast" +version = "1.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8498c871161e1742aaa9d52551b2d6ebdd4c3d45a3be423e3728f33b955be550" +dependencies = [ + "cfg-if", +] + [[package]] name = "criterion" version = "0.8.2" @@ -763,6 +909,15 @@ dependencies = [ "itertools 0.13.0", ] +[[package]] +name = "crossbeam-channel" +version = "0.5.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "98b0cc327b5bc766e7fda9c9260cc0fa81b43a8e240440422dff70788e3f9ef1" +dependencies = [ + "crossbeam-utils", +] + [[package]] name = "crossbeam-deque" version = "0.8.7" @@ -878,11 +1033,20 @@ version = "2.11.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a4ae5f15dda3c708c0ade84bfee31ccab44a3da4f88015ed22f63732abe300c8" +[[package]] +name = "data-url" +version = "0.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "be1e0bca6c3637f992fc1cc7cbc52a78c1ef6db076dbf1059c4323d6a2048376" + [[package]] name = "deranged" version = "0.5.8" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7cd812cc2bc1d69d4764bd80df88b4317eaef9e773c75226407d9bc0876b211c" +dependencies = [ + "serde_core", +] [[package]] name = "derive_builder" @@ -954,6 +1118,12 @@ version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "92773504d58c093f6de2459af4af33faa518c13451eb8f2b5698ed3d36e7c813" +[[package]] +name = "dyn-clone" +version = "1.0.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d0881ea181b1df73ff77ffaaf9c7544ecc11e82fba9b5f27b262a3c73a332555" + [[package]] name = "either" version = "1.16.0" @@ -966,12 +1136,42 @@ version = "1.0.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + [[package]] name = "esaxx-rs" version = "0.1.10" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d817e038c30374a4bcb22f94d0a8a0e216958d4c3dcde369b1439fec4bdda6e6" +[[package]] +name = "event-listener" +version = "5.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a23add41df1562121a9393cb065eab5146a1242410f23a644851e90cfd669d2" +dependencies = [ + "parking", + "pin-project-lite", +] + +[[package]] +name = "event-listener-strategy" +version = "0.5.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8be9f3dfaaffdae2972880079a491a1a8bb7cbed0b8dd7a347f668b4150a3b93" +dependencies = [ + "event-listener", + "pin-project-lite", +] + [[package]] name = "fastrand" version = "2.5.0" @@ -984,6 +1184,17 @@ version = "0.1.9" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" +[[package]] +name = "flate2" +version = "1.1.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e634e2e0ebac1ee034020da1ca582e17ffe4e0f5e985823721e168928136dcb" +dependencies = [ + "crc32fast", + "miniz_oxide", + "zlib-rs", +] + [[package]] name = "fnv" version = "1.0.7" @@ -1005,6 +1216,21 @@ version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" +[[package]] +name = "futures" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a88cf1f829d945f548cf8fec32c61b1f202b6d93b45848602fc02af4b12ad218" +dependencies = [ + "futures-channel", + "futures-core", + "futures-executor", + "futures-io", + "futures-sink", + "futures-task", + "futures-util", +] + [[package]] name = "futures-channel" version = "0.3.33" @@ -1021,6 +1247,17 @@ version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2cd50c473c80f6d7c3670a752354b8e569b1a7cbfdc0419ec88e5edad85e0dc7" +[[package]] +name = "futures-executor" +version = "0.3.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6754879cc9f2c66f88c6e5c35344bb0bdb0708b0352b1201815667c7eabc7458" +dependencies = [ + "futures-core", + "futures-task", + "futures-util", +] + [[package]] name = "futures-io" version = "0.3.33" @@ -1062,6 +1299,7 @@ version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "a77a90a256fce34da66415271e30f94ee91c57b04b8a2c042d9cf3220179deaa" dependencies = [ + "futures-channel", "futures-core", "futures-io", "futures-macro", @@ -1072,6 +1310,33 @@ dependencies = [ "slab", ] +[[package]] +name = "gcp_auth" +version = "0.12.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "26d27dbcc645b60b8e7f6e2868a9d7102ece97d1bb49c1288b5321fcc67f7260" +dependencies = [ + "async-trait", + "base64 0.22.1", + "bytes", + "chrono", + "http 1.4.2", + "http-body-util", + "hyper 1.10.1", + "hyper-rustls 0.27.9", + "hyper-util", + "ring", + "rustls 0.23.42", + "rustls-pki-types", + "serde", + "serde_json", + "thiserror 2.0.19", + "tokio", + "tracing", + "tracing-futures", + "url", +] + [[package]] name = "generic-array" version = "0.14.7" @@ -1380,6 +1645,30 @@ dependencies = [ "tracing", ] +[[package]] +name = "iana-time-zone" +version = "0.1.65" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e31bc9ad994ba00e440a8aa5c9ef0ec67d5cb5e5cb0cc7f8b744a35b389cc470" +dependencies = [ + "android_system_properties", + "core-foundation-sys", + "iana-time-zone-haiku", + "js-sys", + "log", + "wasm-bindgen", + "windows-core", +] + +[[package]] +name = "iana-time-zone-haiku" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f31827a206f56af32e590ba56d5d2d085f558508192593743f16b2306495269f" +dependencies = [ + "cc", +] + [[package]] name = "icu_collections" version = "2.2.0" @@ -1531,6 +1820,55 @@ version = "1.0.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" +[[package]] +name = "jni" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5efd9a482cf3a427f00d6b35f14332adc7902ce91efb778580e180ff90fa3498" +dependencies = [ + "cfg-if", + "combine", + "jni-macros", + "jni-sys", + "log", + "simd_cesu8", + "thiserror 2.0.19", + "walkdir", + "windows-link", +] + +[[package]] +name = "jni-macros" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a00109accc170f0bdb141fed3e393c565b6f5e072365c3bd58f5b062591560a3" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "simd_cesu8", + "syn 2.0.119", +] + +[[package]] +name = "jni-sys" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6377a88cb3910bee9b0fa88d4f42e1d2da8e79915598f65fb0c7ee14c878af2" +dependencies = [ + "jni-sys-macros", +] + +[[package]] +name = "jni-sys-macros" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "38c0b942f458fe50cdac086d2f946512305e5631e720728f2a61aabcd47a6264" +dependencies = [ + "quote", + "syn 2.0.119", +] + [[package]] name = "jobserver" version = "0.1.35" @@ -1574,7 +1912,7 @@ dependencies = [ "futures-util", "litellm-config", "litellm-core", - "reqwest", + "reqwest 0.12.28", "rustls 0.23.42", "rustls-native-certs", "serde", @@ -1607,19 +1945,27 @@ dependencies = [ "aws-sigv4", "aws-smithy-runtime-api", "aws-types", + "azure_core", + "azure_identity", "base64 0.22.1", + "data-url", + "gcp_auth", + "moka", "rand 0.8.7", - "reqwest", + "reqwest 0.12.28", "rstest", "serde", "serde_json", "serde_path_to_error", "sha2 0.10.9", + "strum", + "subtle", "thiserror 2.0.19", "tokio", "tracing", "tracing-subscriber", "url", + "veil", ] [[package]] @@ -1674,6 +2020,15 @@ version = "0.8.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0" +[[package]] +name = "lock_api" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" +dependencies = [ + "scopeguard", +] + [[package]] name = "log" version = "0.4.33" @@ -1736,6 +2091,16 @@ version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "68354c5c6bd36d73ff3feceb05efa59b6acb7626617f4962be322a825e61f79a" +[[package]] +name = "miniz_oxide" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b63fbc4a50860e98e7b2aa7804ded1db5cbc3aff9193adaff57a6931bf7c4b4c" +dependencies = [ + "adler2", + "simd-adler32", +] + [[package]] name = "mio" version = "1.2.2" @@ -1747,6 +2112,26 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "moka" +version = "0.12.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4293f18e7567a1caf3c584855554377025c65e0aa445344d04171f5ad63d19b9" +dependencies = [ + "async-lock", + "crossbeam-channel", + "crossbeam-epoch", + "crossbeam-utils", + "equivalent", + "event-listener", + "futures-util", + "parking_lot", + "portable-atomic", + "smallvec", + "tagptr", + "uuid", +] + [[package]] name = "monostate" version = "0.1.18" @@ -1859,6 +2244,35 @@ dependencies = [ "winapi", ] +[[package]] +name = "parking" +version = "2.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f38d5652c16fde515bb1ecef450ab0f6a219d619a7274976324d5e377f7dceba" + +[[package]] +name = "parking_lot" +version = "0.12.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93857453250e3077bd71ff98b6a65ea6621a19bb0f559a85248955ac12c45a1a" +dependencies = [ + "lock_api", + "parking_lot_core", +] + +[[package]] +name = "parking_lot_core" +version = "0.9.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1" +dependencies = [ + "cfg-if", + "libc", + "redox_syscall", + "smallvec", + "windows-link", +] + [[package]] name = "paste" version = "1.0.15" @@ -1877,6 +2291,26 @@ version = "2.3.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" +[[package]] +name = "pin-project" +version = "1.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2466b2336ed02bcdca6b294417127b90ec92038d1d5c4fbeac971a922e0e0924" +dependencies = [ + "pin-project-internal", +] + +[[package]] +name = "pin-project-internal" +version = "1.1.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c96395f0a926bc13b1c17622aaddda1ecb55d49c8f1bf9777e4d877800a43f8b" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "pin-project-lite" version = "0.2.17" @@ -2078,6 +2512,7 @@ version = "0.11.16" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2f4bfc015262b9df63c8845072ce59068853ff5872180c2ce2f13038b970e560" dependencies = [ + "aws-lc-rs", "bytes", "getrandom 0.4.3", "lru-slab", @@ -2245,6 +2680,15 @@ dependencies = [ "crossbeam-utils", ] +[[package]] +name = "redox_syscall" +version = "0.5.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" +dependencies = [ + "bitflags", +] + [[package]] name = "regex" version = "1.13.1" @@ -2325,11 +2769,49 @@ dependencies = [ "url", "wasm-bindgen", "wasm-bindgen-futures", - "wasm-streams", + "wasm-streams 0.4.2", "web-sys", "webpki-roots", ] +[[package]] +name = "reqwest" +version = "0.13.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "16a1cfa75cc186dd73d5818e510e042e40927bccc9c236b061cea97e1eb08029" +dependencies = [ + "base64 0.23.1", + "bytes", + "futures-core", + "futures-util", + "http 1.4.2", + "http-body 1.1.0", + "http-body-util", + "hyper 1.10.1", + "hyper-rustls 0.27.9", + "hyper-util", + "js-sys", + "log", + "percent-encoding", + "pin-project-lite", + "quinn", + "rustls 0.23.42", + "rustls-pki-types", + "rustls-platform-verifier", + "sync_wrapper", + "tokio", + "tokio-rustls 0.26.4", + "tokio-util", + "tower", + "tower-http", + "tower-service", + "url", + "wasm-bindgen", + "wasm-bindgen-futures", + "wasm-streams 0.5.0", + "web-sys", +] + [[package]] name = "ring" version = "0.17.14" @@ -2437,6 +2919,33 @@ dependencies = [ "zeroize", ] +[[package]] +name = "rustls-platform-verifier" +version = "0.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "26d1e2536ce4f35f4846aa13bff16bd0ff40157cdb14cc056c7b14ba41233ba0" +dependencies = [ + "core-foundation", + "core-foundation-sys", + "jni", + "log", + "once_cell", + "rustls 0.23.42", + "rustls-native-certs", + "rustls-platform-verifier-android", + "rustls-webpki 0.103.13", + "security-framework", + "security-framework-sys", + "webpki-root-certs", + "windows-sys 0.61.2", +] + +[[package]] +name = "rustls-platform-verifier-android" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f87165f0995f63a9fbeea62b64d10b4d9d8e78ec6d7d51fb2125fda7bb36788f" + [[package]] name = "rustls-webpki" version = "0.101.7" @@ -2489,6 +2998,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "scopeguard" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" + [[package]] name = "sct" version = "0.7.1" @@ -2642,6 +3157,38 @@ version = "2.0.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba" +[[package]] +name = "signal-hook-registry" +version = "1.4.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4db69cba1110affc0e9f7bcd48bbf87b3f4fc7c61fc9155afd4c469eb3d6c1b" +dependencies = [ + "errno", + "libc", +] + +[[package]] +name = "simd-adler32" +version = "0.3.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a219298ac11a56ea9a6d2120044824d6f01aeb034955e7af7bc16858527deea" + +[[package]] +name = "simd_cesu8" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11031e251abf8611c80f460e19dbdeb54a66db918e49c65a7065b46ac7aec520" +dependencies = [ + "rustc_version", + "simdutf8", +] + +[[package]] +name = "simdutf8" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3a9fe34e3e7a50316060351f37187a3f546bce95496156754b601a5fa71b76e" + [[package]] name = "slab" version = "0.4.12" @@ -2704,6 +3251,27 @@ version = "0.11.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" +[[package]] +name = "strum" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9628de9b8791db39ceda2b119bbe13134770b56c138ec1d3af810d045c04f9bd" +dependencies = [ + "strum_macros", +] + +[[package]] +name = "strum_macros" +version = "0.28.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ab85eea0270ee17587ed4156089e10b9e6880ee688791d45a905f5b1ca36f664" +dependencies = [ + "heck", + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "subtle" version = "2.6.1" @@ -2752,6 +3320,12 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "tagptr" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b2093cf4c8eb1e67749a6762251bc9cd836b6fc171623bd0a9d324d37af2417" + [[package]] name = "target-lexicon" version = "0.13.5" @@ -2915,6 +3489,7 @@ dependencies = [ "libc", "mio", "pin-project-lite", + "signal-hook-registry", "socket2 0.6.5", "tokio-macros", "windows-sys 0.61.2", @@ -3032,12 +3607,17 @@ version = "0.6.11" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" dependencies = [ + "async-compression", "bitflags", "bytes", + "futures-core", "futures-util", "http 1.4.2", "http-body 1.1.0", + "http-body-util", "pin-project-lite", + "tokio", + "tokio-util", "tower", "tower-layer", "tower-service", @@ -3088,6 +3668,16 @@ dependencies = [ "once_cell", ] +[[package]] +name = "tracing-futures" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "97d095ae15e245a057c8e8451bab9b3ee1e1f68e9ba2b4fbc18d0ac5237835f2" +dependencies = [ + "pin-project", + "tracing", +] + [[package]] name = "tracing-subscriber" version = "0.3.23" @@ -3131,6 +3721,57 @@ version = "1.20.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" +[[package]] +name = "typespec" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "753a2fe021e407d4fc9ee6f4f0a33403cc306d5c54c4e4ebe1b8cbde0ca052b9" +dependencies = [ + "base64 0.22.1", + "bytes", + "futures", + "serde", + "serde_json", + "url", +] + +[[package]] +name = "typespec_client_core" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0373af0f9d4f580b3a1a9d9639cedaabe015ed262b35bfbe13941bfb14fe1ea6" +dependencies = [ + "async-trait", + "base64 0.22.1", + "bytes", + "dyn-clone", + "futures", + "pin-project", + "rand 0.10.2", + "reqwest 0.13.5", + "serde", + "serde_json", + "time", + "tokio", + "tracing", + "typespec", + "typespec_macros", + "url", + "uuid", +] + +[[package]] +name = "typespec_macros" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2c608f4427943f8adb211abc95c87672b1b98847152783507d54e3246e502f60" +dependencies = [ + "proc-macro2", + "quote", + "rustc_version", + "syn 2.0.119", +] + [[package]] name = "unicase" version = "2.9.0" @@ -3206,10 +3847,32 @@ version = "1.24.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "bf3923a6f5c4c6382e0b653c4117f48d631ea17f38ed86e2a828e6f7412f5239" dependencies = [ + "getrandom 0.4.3", "js-sys", "wasm-bindgen", ] +[[package]] +name = "veil" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7352f0bbf3ab98911b0c0277065094c1b1ec79bbc85fa3b7d16bf1859c3d96f" +dependencies = [ + "once_cell", + "veil-macros", +] + +[[package]] +name = "veil-macros" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47a3f4f06d904eb789b935253752ba6bcc1dfa61349f8d5341c66abe070b44e5" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "version_check" version = "0.9.5" @@ -3324,6 +3987,19 @@ dependencies = [ "web-sys", ] +[[package]] +name = "wasm-streams" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9d1ec4f6517c9e11ae630e200b2b65d193279042e28edd4a2cda233e46670bbb" +dependencies = [ + "futures-util", + "js-sys", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + [[package]] name = "web-sys" version = "0.3.103" @@ -3344,6 +4020,15 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "webpki-root-certs" +version = "1.0.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b96554aa2acc8ccdb7e1c9a58a7a68dd5d13bccc69cd124cb09406db612a1c9b" +dependencies = [ + "rustls-pki-types", +] + [[package]] name = "webpki-roots" version = "1.0.9" @@ -3384,12 +4069,65 @@ version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "712e227841d057c1ee1cd2fb22fa7e5a5461ae8e48fa2ca79ec42cfc1931183f" +[[package]] +name = "windows-core" +version = "0.62.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8e83a14d34d0623b51dce9581199302a221863196a1dde71a7663a4c2be9deb" +dependencies = [ + "windows-implement", + "windows-interface", + "windows-link", + "windows-result", + "windows-strings", +] + +[[package]] +name = "windows-implement" +version = "0.60.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "053e2e040ab57b9dc951b72c264860db7eb3b0200ba345b4e4c3b14f67855ddf" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + +[[package]] +name = "windows-interface" +version = "0.59.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f316c4a2570ba26bbec722032c4099d8c8bc095efccdc15688708623367e358" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "windows-link" version = "0.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" +[[package]] +name = "windows-result" +version = "0.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7781fa89eaf60850ac3d2da7af8e5242a5ea78d1a11c49bf2910bb5a73853eb5" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-strings" +version = "0.5.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7837d08f69c77cf6b07689544538e017c1bfcf57e34b4c0ff58e6c2cd3b37091" +dependencies = [ + "windows-link", +] + [[package]] name = "windows-sys" version = "0.52.0" @@ -3602,6 +4340,12 @@ dependencies = [ "syn 2.0.119", ] +[[package]] +name = "zlib-rs" +version = "0.6.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "34b31d188d9d685a4f9c7b46d6e36631b07058d2cfe190267adce54dc230bf12" + [[package]] name = "zmij" version = "1.0.23" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index f3e54e5b2aa..5f25e69a1f8 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -41,8 +41,14 @@ 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"] } base64 = "0.22" +gcp_auth = "0.12.7" +azure_core = "1.0.0" +azure_identity = { version = "1.0.0", features = ["tokio"] } +moka = { version = "0.12.16", features = ["future"] } +strum = { version = "0.28.0", features = ["derive"] } url = "2.5.8" criterion = "0.8.2" +veil = "0.3.0" [profile.release] opt-level = 3 diff --git a/litellm-rust/crates/ai-gateway/Dockerfile b/litellm-rust/crates/ai-gateway/Dockerfile index 2bc3c05ad7e..72ac25ce1d6 100644 --- a/litellm-rust/crates/ai-gateway/Dockerfile +++ b/litellm-rust/crates/ai-gateway/Dockerfile @@ -14,15 +14,20 @@ # ---- Chef ------------------------------------------------------------------- # cargo-chef caches the dependency build so only the gateway crate recompiles on # a source-only change. python3-dev is present in every rust stage because the -# `python-config` feature links libpython via pyo3 (even in the cook step). -FROM rust:1.90-slim-bookworm AS chef +# `python-config` feature links libpython via pyo3 (even in the cook step), and +# python3-pip builds the litellm wheel in the builder stage. +FROM rust:1.98-slim-bookworm AS chef ENV PYO3_PYTHON=python3.11 +# rustup reads rust-toolchain.toml from any parent of the working directory, so +# copying it in is what keeps every cargo call below on the repo's pinned +# channel rather than on whatever the base image happens to ship. +COPY rust-toolchain.toml /build/rust-toolchain.toml +WORKDIR /build/litellm-rust RUN apt-get update \ && apt-get install -y --no-install-recommends \ - python3 python3-dev pkg-config libssl-dev clang \ + python3 python3-dev python3-pip pkg-config libssl-dev clang \ && rm -rf /var/lib/apt/lists/* \ && cargo install cargo-chef --locked --version 0.1.77 -WORKDIR /build/litellm-rust # ---- Planner ---------------------------------------------------------------- # Produce the dependency recipe from the rust workspace manifests + Cargo.lock. @@ -43,6 +48,19 @@ RUN cargo chef cook --locked --release \ COPY litellm-rust/ . RUN cargo build --locked --release -p litellm-ai-gateway --bin litellm-ai-gateway --features server,python-config +# The root pyproject builds with maturin against litellm-rust/crates/python-bridge, +# so the wheel is built here, next to the crate sources and the cargo toolchain, +# and the runtime stage installs the artifact instead of compiling anything. +# litellm[proxy] pins litellm-enterprise and litellm-proxy-extras to the versions +# in this repo, and those hit PyPI hours after every version bump merges, so both +# wheels are built from the repo too instead of being resolved from PyPI. +COPY pyproject.toml README.md LICENSE /build/ +COPY litellm/ /build/litellm/ +COPY enterprise/ /build/enterprise/ +COPY litellm-proxy-extras/ /build/litellm-proxy-extras/ +RUN pip3 wheel --no-cache-dir --no-deps --wheel-dir /build/dist \ + /build /build/enterprise /build/litellm-proxy-extras + # ---- Runtime ---------------------------------------------------------------- # python:3.11-slim-bookworm ships libpython3.11, matching the builder's PyO3 # 3.11 ABI so the embedded interpreter links and imports cleanly. @@ -56,11 +74,16 @@ RUN apt-get update \ WORKDIR /app # Install litellm (with proxy extras) FROM THIS REPO'S SOURCE so -# `import litellm.proxy.read_model_list` works — it is not on PyPI yet. Copy the -# package + packaging metadata, then pip install the proxy extra. -COPY pyproject.toml README.md LICENSE ./ -COPY litellm/ ./litellm/ -RUN pip install --no-cache-dir ".[proxy]" +# `import litellm.proxy.read_model_list` works — it is not on PyPI yet. The two +# sibling wheels come from the builder as well, so the pins in litellm[proxy] +# resolve against them and never wait on a PyPI publish. +COPY --from=builder /build/dist/*.whl /tmp/wheels/ +RUN wheel="$(ls /tmp/wheels/litellm-*.whl)" \ + && pip install --no-cache-dir \ + /tmp/wheels/litellm_enterprise-*.whl \ + /tmp/wheels/litellm_proxy_extras-*.whl \ + "${wheel}[proxy]" \ + && rm -rf /tmp/wheels # The compiled gateway binary (pure-Rust realtime hot path; Python is load-time # only). diff --git a/litellm-rust/crates/ai-gateway/Dockerfile.dockerignore b/litellm-rust/crates/ai-gateway/Dockerfile.dockerignore index 030ee6a37c5..d1386ff684d 100644 --- a/litellm-rust/crates/ai-gateway/Dockerfile.dockerignore +++ b/litellm-rust/crates/ai-gateway/Dockerfile.dockerignore @@ -9,19 +9,28 @@ # Strategy: ignore everything, then re-include only what the build needs: # - litellm/ (pip install . needs the full package + proxy reader) # - litellm-rust/ (the rust workspace; Cargo.lock + crate sources) -# - pyproject.toml / README.md / LICENSE (packaging metadata for pip install) +# - enterprise/ (litellm/proxy/enterprise symlinks into it; maturin walks it) +# - litellm-proxy-extras/ (built into a wheel alongside enterprise/ for litellm[proxy]) +# - pyproject.toml / README.md / LICENSE (packaging metadata for the wheel build) +# - rust-toolchain.toml (the pinned channel every cargo call in the build uses) * # --- re-include the build inputs --- !litellm/ !litellm-rust/ +!enterprise/ +!litellm-proxy-extras/ !pyproject.toml +!rust-toolchain.toml !README.md !LICENSE # --- prune heavy / irrelevant subpaths back out of the re-included trees --- # Rust build artifacts (huge; regenerated in the builder). **/target/ +# Committed python distribution artifacts; the wheel build does not read them. +enterprise/dist/ +litellm-proxy-extras/dist/ # Python caches and compiled bytecode. **/__pycache__/ **/*.pyc diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs index 098ce071efc..6f48f38c9f6 100644 --- a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -269,7 +269,11 @@ fn guardrail_error_to_core_error(error: GuardrailError) -> Error { fn core_error_kind(error: &Error) -> &'static str { match error { - Error::Auth(_) | Error::MissingApiKey { .. } => "AuthError", + Error::Auth(_) + | Error::MissingApiKey { .. } + | Error::MissingAzureAiCredentials + | Error::MissingAzureDocumentIntelligenceCredentials + | Error::MissingReductoApiKey => "AuthError", Error::InvalidProvider(_) => "InvalidProvider", Error::InvalidRequest(_) => "InvalidRequest", Error::InvalidType { .. } => "InvalidType", diff --git a/litellm-rust/crates/ai-gateway/src/io/realtime.rs b/litellm-rust/crates/ai-gateway/src/io/realtime.rs index 207c31dffa0..1aa31adcc38 100644 --- a/litellm-rust/crates/ai-gateway/src/io/realtime.rs +++ b/litellm-rust/crates/ai-gateway/src/io/realtime.rs @@ -15,6 +15,8 @@ use std::time::Duration; use futures_util::stream::{SplitSink, SplitStream}; use futures_util::{Sink, SinkExt, Stream, StreamExt}; +use litellm_core::AuthError; +use litellm_core::auth::error::MissingCredential; use litellm_core::error::Error; use litellm_core::realtime::transformation::RealtimeProviderConfig; use litellm_core::realtime::types::RealtimeEvent; @@ -32,8 +34,6 @@ use crate::io::tls::connect_upstream; /// Environment variable holding the OpenAI API key (last-resort fallback). const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY"; -const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a realtime call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"; - /// Default **idle** timeout: if neither side sends a frame for this long, the /// session is reaped. It resets on any activity, so it does not cap a healthy /// (continuously streaming) session — it only frees a stalled one (e.g. a @@ -59,7 +59,7 @@ pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result { .ok() .filter(|key| !key.trim().is_empty()) }) - .ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string())) + .ok_or_else(|| Error::from(AuthError::from(MissingCredential::OpenAiRealtimeApiKey))) } /// Open the upstream WebSocket to OpenAI for `(model, api_key, api_base)`. diff --git a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs index 9df3d0c6cc5..7f3b6b0650f 100644 --- a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs +++ b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs @@ -4,7 +4,9 @@ use std::time::Duration; use futures_util::stream::{SplitSink, SplitStream}; use futures_util::{Sink, SinkExt, Stream, StreamExt}; +use litellm_core::AuthError; use litellm_core::Error; +use litellm_core::auth::error::MissingCredential; use litellm_core::providers::openai::responses::transformation::OPENAI_RESPONSES_WS_CONFIG; use litellm_core::responses::types::ResponsesWsEvent; use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig; @@ -23,8 +25,6 @@ use crate::constants::{ }; const OPENAI_API_KEY_ENV: &str = "OPENAI_API_KEY"; -const MISSING_KEY_MESSAGE: &str = "Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable"; - pub type ResponsesUpstreamWs = WebSocketStream>; type UpstreamTx = SplitSink; type UpstreamRx = SplitStream; @@ -120,7 +120,7 @@ pub(crate) fn resolve_api_key(api_key: Option<&str>) -> Result { .ok() .filter(|value| !value.trim().is_empty()) }) - .ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string())) + .ok_or_else(|| Error::from(AuthError::from(MissingCredential::OpenAiResponsesApiKey))) } async fn dial_upstream( diff --git a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs b/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs deleted file mode 100644 index d2be17260a3..00000000000 --- a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs +++ /dev/null @@ -1,529 +0,0 @@ -use std::net::IpAddr; -use std::time::{Duration, Instant}; - -use base64::Engine; -use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; -use litellm_core::error::Error; -use litellm_core::ocr::transformation::OcrProviderConfig; -use reqwest::Url; -use serde_json::{Map, Value}; - -use litellm_core::providers::azure_ai::ocr::transformation::{ - AZURE_AI_OCR_CONFIG, AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG, -}; -use litellm_core::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; -use litellm_core::providers::reducto::ocr::transformation as reducto; -use litellm_core::providers::vertex_ai::ocr::transformation as vertex_ai; -use litellm_core::providers::vertex_ai::ocr::transformation::{ - VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG, -}; - -use crate::client::http_client; - -const ERROR_BODY_MAX_CHARS: usize = 256; -const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120; -const DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB: f64 = 50.0; -const MAX_SAFE_FETCH_REDIRECTS: usize = 10; - -pub(super) fn truncate_error_body(body: &str) -> String { - if body.chars().count() <= ERROR_BODY_MAX_CHARS { - return body.to_string(); - } - let truncated: String = body.chars().take(ERROR_BODY_MAX_CHARS).collect(); - format!("{truncated}... (truncated)") -} - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(super) fn ocr_provider_config( - provider: &str, - model: &str, -) -> Option<&'static dyn OcrProviderConfig> { - match provider { - "mistral" => Some(&MISTRAL_OCR_CONFIG), - "reducto" => reducto::config_for_model(model), - "azure_ai" if is_azure_document_intelligence_model(model) => { - Some(&AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG) - } - "azure_ai" => Some(&AZURE_AI_OCR_CONFIG), - "vertex_ai" if vertex_ai::is_deepseek_model(model) => Some(&VERTEX_AI_DEEPSEEK_OCR_CONFIG), - "vertex_ai" => Some(&VERTEX_AI_OCR_CONFIG), - _ => None, - } -} - -fn is_azure_document_intelligence_model(model: &str) -> bool { - let model = model.to_ascii_lowercase(); - model.contains("doc-intelligence") || model.contains("documentintelligence") -} - -pub(super) fn string_headers( - extra_headers: Option>, -) -> Result, Error> { - extra_headers - .unwrap_or_default() - .into_iter() - .map(|(key, value)| { - value - .as_str() - .map(|value| (key.clone(), value.to_string())) - .ok_or_else(|| { - Error::InvalidRequest(format!( - "OCR extra_headers.{key} must be a string, got {}", - litellm_core::error::json_type_name(&value) - )) - }) - }) - .collect() -} - -fn document_url_field(document: &Value) -> Result, Error> { - let Some(object) = document.as_object() else { - return Ok(None); - }; - let Some(doc_type) = object.get("type").and_then(Value::as_str) else { - return Ok(None); - }; - let field = match doc_type { - "document_url" => "document_url", - "image_url" => "image_url", - _ => return Ok(None), - }; - let Some(url) = object.get(field).and_then(Value::as_str) else { - return Ok(None); - }; - Ok(Some((field, url))) -} - -fn is_url_requiring_fetch(url: &str) -> bool { - !url.starts_with("data:") && (url.starts_with("http://") || url.starts_with("https://")) -} - -fn max_document_download_bytes() -> u64 { - let max_size_mb = std::env::var("MAX_IMAGE_URL_DOWNLOAD_SIZE_MB") - .ok() - .and_then(|value| value.parse::().ok()) - .unwrap_or(DEFAULT_MAX_IMAGE_URL_DOWNLOAD_SIZE_MB); - (max_size_mb.max(0.0) * 1024.0 * 1024.0) as u64 -} - -fn is_blocked_ip(ip: IpAddr) -> bool { - match ip { - IpAddr::V4(ip) => { - ip.is_private() - || ip.is_loopback() - || ip.is_link_local() - || ip.is_broadcast() - || ip.is_multicast() - || ip.is_unspecified() - } - IpAddr::V6(ip) => { - let first_segment = ip.segments()[0]; - let is_unique_local = (first_segment & 0xfe00) == 0xfc00; - let is_link_local = (first_segment & 0xffc0) == 0xfe80; - ip.is_loopback() - || ip.is_unspecified() - || ip.is_multicast() - || is_unique_local - || is_link_local - || ip - .to_ipv4_mapped() - .or_else(|| ip.to_ipv4()) - .map(|v4| is_blocked_ip(IpAddr::V4(v4))) - .unwrap_or(false) - } - } -} - -fn blocked_url_error(url: &Url) -> Error { - Error::InvalidRequest(format!( - "OCR document URL rejected by SSRF protection: {url}" - )) -} - -async fn validate_safe_fetch_url(url: &Url) -> Result<(), Error> { - if !matches!(url.scheme(), "http" | "https") { - return Err(blocked_url_error(url)); - } - - let host = url.host_str().ok_or_else(|| blocked_url_error(url))?; - if let Ok(ip) = host.parse::() { - if is_blocked_ip(ip) { - return Err(blocked_url_error(url)); - } - return Ok(()); - } - - let port = url - .port_or_known_default() - .ok_or_else(|| blocked_url_error(url))?; - let addresses = tokio::net::lookup_host((host, port)) - .await - .map_err(|err| Error::Network(err.to_string()))?; - let mut saw_address = false; - for address in addresses { - saw_address = true; - if is_blocked_ip(address.ip()) { - return Err(blocked_url_error(url)); - } - } - if !saw_address { - return Err(blocked_url_error(url)); - } - Ok(()) -} - -fn redirect_location(response: &reqwest::Response, url: &Url) -> Result { - let location = response - .headers() - .get(reqwest::header::LOCATION) - .and_then(|value| value.to_str().ok()) - .ok_or_else(|| { - Error::InvalidResponse("OCR document redirect missing Location header".to_string()) - })?; - url.join(location) - .map_err(|err| Error::InvalidResponse(format!("invalid OCR document redirect: {err}"))) -} - -async fn safe_get_document_url(url: &str) -> Result<(Url, reqwest::Response), Error> { - let client = reqwest::Client::builder() - .redirect(reqwest::redirect::Policy::none()) - .build() - .map_err(|err| Error::Network(err.to_string()))?; - let mut current_url = Url::parse(url) - .map_err(|err| Error::InvalidRequest(format!("invalid OCR document URL: {err}")))?; - - for _ in 0..MAX_SAFE_FETCH_REDIRECTS { - validate_safe_fetch_url(¤t_url).await?; - let response = client - .get(current_url.clone()) - .send() - .await - .map_err(|err| Error::Network(err.to_string()))?; - if !response.status().is_redirection() { - return Ok((current_url, response)); - } - current_url = redirect_location(&response, ¤t_url)?; - } - - Err(Error::InvalidRequest( - "Too many redirects while fetching OCR document URL".to_string(), - )) -} - -fn enforce_download_size(content_length: u64, max_bytes: u64, url: &Url) -> Result<(), Error> { - if max_bytes == 0 { - return Err(Error::InvalidRequest(format!( - "OCR document URL download is disabled (MAX_IMAGE_URL_DOWNLOAD_SIZE_MB=0). url={url}" - ))); - } - if content_length > max_bytes { - let size_mb = content_length as f64 / (1024.0 * 1024.0); - let max_size_mb = max_bytes as f64 / (1024.0 * 1024.0); - return Err(Error::InvalidRequest(format!( - "OCR document size ({size_mb:.2}MB) exceeds maximum allowed size ({max_size_mb:.2}MB). url={url}" - ))); - } - Ok(()) -} - -async fn read_response_with_limit( - mut response: reqwest::Response, - url: &Url, -) -> Result, Error> { - let max_bytes = max_document_download_bytes(); - if let Some(content_length) = response.content_length() { - enforce_download_size(content_length, max_bytes, url)?; - } else { - enforce_download_size(0, max_bytes, url)?; - } - - let mut bytes = Vec::new(); - let mut bytes_downloaded: u64 = 0; - while let Some(chunk) = response - .chunk() - .await - .map_err(|err| Error::Network(err.to_string()))? - { - bytes_downloaded += chunk.len() as u64; - enforce_download_size(bytes_downloaded, max_bytes, url)?; - bytes.extend_from_slice(&chunk); - } - Ok(bytes) -} - -pub(super) async fn convert_document_url_to_data_uri(document: Value) -> Result { - let Some((field, url)) = document_url_field(&document)? else { - return Ok(document); - }; - if !is_url_requiring_fetch(url) { - return Ok(document); - } - - let (final_url, response) = safe_get_document_url(url).await?; - let status = response.status(); - if !status.is_success() { - let body = response.text().await.unwrap_or_default(); - return Err(Error::Http { - status: status.as_u16(), - body: truncate_error_body(&body), - }); - } - let content_type = response - .headers() - .get(reqwest::header::CONTENT_TYPE) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.split(';').next()) - .map(str::trim) - .filter(|value| !value.is_empty()) - .unwrap_or("application/octet-stream") - .to_string(); - let bytes = read_response_with_limit(response, &final_url).await?; - let data_uri = format!( - "data:{content_type};base64,{}", - BASE64_STANDARD.encode(bytes) - ); - - let mut transformed = document - .as_object() - .cloned() - .ok_or_else(|| Error::InvalidRequest("OCR document must be an object".to_string()))?; - transformed.insert(field.to_string(), Value::String(data_uri)); - Ok(Value::Object(transformed)) -} - -fn same_origin(left: &str, right: &str) -> bool { - let Ok(left) = reqwest::Url::parse(left) else { - return false; - }; - let Ok(right) = reqwest::Url::parse(right) else { - return false; - }; - left.scheme() == right.scheme() - && left.host_str() == right.host_str() - && left.port_or_known_default() == right.port_or_known_default() -} - -fn retry_after_secs(response: &reqwest::Response) -> u64 { - response - .headers() - .get(reqwest::header::RETRY_AFTER) - .and_then(|value| value.to_str().ok()) - .and_then(|value| value.parse::().ok()) - .unwrap_or(2) -} - -fn operation_status(response_json: &Value) -> Result<&str, Error> { - let status = response_json - .get("status") - .and_then(Value::as_str) - .ok_or(Error::MissingField("status"))?; - match status { - "succeeded" => Ok("succeeded"), - "running" | "notStarted" => Ok("running"), - "failed" => { - let message = response_json - .get("error") - .and_then(|error| error.get("message")) - .and_then(Value::as_str) - .unwrap_or("Unknown error"); - Err(Error::InvalidResponse(format!( - "Azure Document Intelligence analysis failed: {message}" - ))) - } - other => Err(Error::InvalidResponse(format!( - "Unknown operation status: {other}" - ))), - } -} - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(super) async fn poll_document_intelligence( - operation_url: &str, - original_url: &str, - headers: &[(String, String)], - timeout: Option, -) -> Result { - if !same_origin(operation_url, original_url) { - return Err(Error::InvalidResponse( - "Azure Document Intelligence: rejected cross-origin polling URL".to_string(), - )); - } - - let start = Instant::now(); - let timeout = timeout.unwrap_or(Duration::from_secs( - AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS, - )); - loop { - if start.elapsed() > timeout { - return Err(Error::Network(format!( - "Azure Document Intelligence operation polling timed out after {} seconds", - timeout.as_secs() - ))); - } - - let mut request_builder = http_client().get(operation_url); - for (key, value) in headers { - if key.eq_ignore_ascii_case("ocp-apim-subscription-key") { - request_builder = request_builder.header(key, value); - } - } - let response = request_builder - .send() - .await - .map_err(|err| Error::Network(err.to_string()))?; - let retry_after = retry_after_secs(&response); - let status = response.status(); - let text = response - .text() - .await - .map_err(|err| Error::Network(err.to_string()))?; - if !status.is_success() { - return Err(Error::Http { - status: status.as_u16(), - body: truncate_error_body(&text), - }); - } - let response_json: Value = serde_json::from_str(&text).map_err(|err| { - Error::InvalidResponse(format!("invalid Azure DI poll response JSON: {err}")) - })?; - if operation_status(&response_json)? == "succeeded" { - return Ok(response_json); - } - tokio::time::sleep(Duration::from_secs(retry_after)).await; - } -} - -#[cfg(test)] -mod tests { - use litellm_core::ocr::transformation::OcrResponseHandling; - use serde_json::json; - - use super::*; - - #[test] - fn blocks_private_and_metadata_ips() { - assert!(is_blocked_ip("127.0.0.1".parse().unwrap())); - assert!(is_blocked_ip("10.0.0.1".parse().unwrap())); - assert!(is_blocked_ip("169.254.169.254".parse().unwrap())); - assert!(is_blocked_ip("::1".parse().unwrap())); - assert!(is_blocked_ip("fd00::1".parse().unwrap())); - assert!(is_blocked_ip("fe80::1".parse().unwrap())); - assert!(is_blocked_ip("::ffff:169.254.169.254".parse().unwrap())); - assert!(is_blocked_ip("::ffff:10.0.0.1".parse().unwrap())); - assert!(!is_blocked_ip("8.8.8.8".parse().unwrap())); - assert!(!is_blocked_ip("::ffff:8.8.8.8".parse().unwrap())); - } - - #[tokio::test] - async fn convert_document_url_rejects_loopback_fetch() { - let error = convert_document_url_to_data_uri(json!({ - "type": "image_url", - "image_url": "http://127.0.0.1/image.png" - })) - .await - .unwrap_err(); - - assert!(matches!( - error, - Error::InvalidRequest(message) - if message.contains("SSRF protection") - )); - } - - #[tokio::test] - async fn convert_document_url_leaves_data_uri_untouched() { - let document = json!({ - "type": "image_url", - "image_url": "data:image/png;base64,abcd" - }); - - let transformed = convert_document_url_to_data_uri(document.clone()) - .await - .unwrap(); - - assert_eq!(transformed, document); - } - - #[test] - fn truncate_error_body_passes_short_strings_through() { - let body = "Unauthorized"; - assert_eq!(truncate_error_body(body), "Unauthorized"); - } - - #[test] - fn truncate_error_body_caps_long_payloads() { - let body = "x".repeat(306); - let truncated = truncate_error_body(&body); - - assert!(truncated.ends_with("... (truncated)")); - let prefix_chars = truncated - .strip_suffix("... (truncated)") - .expect("truncated marker present") - .chars() - .count(); - assert_eq!(prefix_chars, 256); - } - - #[test] - fn truncate_error_body_does_not_split_multibyte_chars() { - let body = "é".repeat(266); - let truncated = truncate_error_body(&body); - assert!(truncated.is_char_boundary(truncated.len())); - } - - #[test] - fn ocr_dispatch_supports_migrated_providers() { - assert!(ocr_provider_config("mistral", "mistral-ocr-latest").is_some()); - assert!( - ocr_provider_config("azure_ai", "pixtral-12b-2409") - .expect("azure ai config resolves") - .requires_data_uri_document() - ); - assert_eq!( - ocr_provider_config("azure_ai", "doc-intelligence/prebuilt-read") - .expect("document intelligence config resolves") - .response_handling(), - OcrResponseHandling::AzureDocumentIntelligencePoll - ); - assert!( - ocr_provider_config("vertex_ai", "deepseek-ocr-maas") - .expect("vertex deepseek config resolves") - .supported_ocr_params() - .contains(&"temperature") - ); - assert!(ocr_provider_config("openai", "gpt-4o").is_none()); - } - - #[test] - fn string_headers_accepts_string_values() { - let headers = json!({ - "x-trace-id": "trace-1" - }) - .as_object() - .unwrap() - .clone(); - - assert_eq!( - string_headers(Some(headers)).expect("string headers accepted"), - vec![("x-trace-id".to_string(), "trace-1".to_string())] - ); - } - - #[test] - fn string_headers_rejects_non_string_values() { - let headers = json!({ - "x-retry-count": 3 - }) - .as_object() - .unwrap() - .clone(); - - let err = string_headers(Some(headers)).expect_err("non-string header rejected"); - assert_eq!( - err, - Error::InvalidRequest( - "OCR extra_headers.x-retry-count must be a string, got number".to_string() - ) - ); - } -} diff --git a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs deleted file mode 100644 index 6c6e12724cd..00000000000 --- a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs +++ /dev/null @@ -1,84 +0,0 @@ -use litellm_core::error::Error; -use litellm_core::http_utils::http_request; -use litellm_core::ocr::transformation::OcrResponseHandling; -use serde_json::Value; - -use super::common_utils::{poll_document_intelligence, truncate_error_body}; -use super::hooks::OcrLifecycleHooks; -use super::types::PreparedOcrRequest; -use crate::client::http_client; - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(crate) async fn execute_ocr_provider_call( - request: PreparedOcrRequest, - hooks: &OcrLifecycleHooks, -) -> Result { - let request = hooks.prepare_provider_request(request).await?; - let mut request_builder = http_client().post(&request.url).json(&request.body); - for (key, value) in &request.upstream_headers { - request_builder = request_builder.header(key, value); - } - if let Some(duration) = request.timeout { - request_builder = request_builder.timeout(duration); - } - - let response = http_request(request_builder) - .await - .map_err(|err| Error::Network(err.to_string()))?; - - let status = response.status(); - if request.config.response_handling() == OcrResponseHandling::AzureDocumentIntelligencePoll - && status.as_u16() == 202 - { - let operation_url = response - .headers() - .get("operation-location") - .and_then(|value| value.to_str().ok()) - .map(str::to_string) - .ok_or_else(|| { - Error::InvalidResponse( - "Azure Document Intelligence returned 202 but no Operation-Location header found" - .to_string(), - ) - })?; - let response_json = poll_document_intelligence( - &operation_url, - &request.url, - &request.upstream_headers, - request.timeout, - ) - .await?; - return Ok(request - .config - .transform_ocr_response_with_params( - &request.model, - response_json, - &request.optional_params, - )? - .into_json()); - } - - let text = response - .text() - .await - .map_err(|err| Error::Network(err.to_string()))?; - - if !status.is_success() { - return Err(Error::Http { - status: status.as_u16(), - body: truncate_error_body(&text), - }); - } - - let response_json: Value = serde_json::from_str(&text) - .map_err(|err| Error::InvalidResponse(format!("invalid OCR response JSON: {err}")))?; - - Ok(request - .config - .transform_ocr_response_with_params( - &request.model, - response_json, - &request.optional_params, - )? - .into_json()) -} diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs deleted file mode 100644 index d1dd811f7ab..00000000000 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ /dev/null @@ -1,401 +0,0 @@ -use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; -use litellm_core::error::Error; -use litellm_core::providers::reducto::ocr::transformation::{ - build_upload_request, extract_document_source, extract_upload_file_id, -}; -use serde_json::{Map, Value, json}; -use std::future::Future; -use std::pin::Pin; - -use super::common_utils::{convert_document_url_to_data_uri, string_headers, truncate_error_body}; -use super::types::{PreparedOcrRequest, ProviderOcrRequest}; -use crate::client::http_client; -use crate::integrations::custom_guardrail::{ - CustomGuardrailRunner, GuardrailContext, GuardrailError, GuardrailRequest, -}; -use crate::integrations::custom_logger::{ - CallType, CallbackTiming, CallbackValue, CustomLoggerRunner, LoggingError, ModelCallDetails, -}; -use crate::integrations::types::{ - RequestMetadata, StandardLoggingMetadata, StandardLoggingPayload, -}; - -pub(crate) struct OcrLifecycleHooks { - logger_runner: CustomLoggerRunner, - guardrail_runner: CustomGuardrailRunner, - request_metadata: RequestMetadata, -} - -type OcrFuture<'a, T> = Pin> + Send + 'a>>; -type OcrLogFuture<'a> = Pin + Send + 'a>>; - -impl OcrLifecycleHooks { - pub(crate) fn new( - logger_runner: CustomLoggerRunner, - guardrail_runner: CustomGuardrailRunner, - request_metadata: RequestMetadata, - ) -> Self { - Self { - logger_runner, - guardrail_runner, - request_metadata, - } - } - - async fn run_pre_call_guardrails( - &self, - request: PreparedOcrRequest, - ) -> Result { - if self.guardrail_runner.is_empty() { - return Ok(request); - } - - let context = guardrail_context(&self.request_metadata); - let guardrail_request = GuardrailRequest::new(json!({ - "model": request.model, - "custom_llm_provider": request.custom_llm_provider, - "document": request.document, - "optional_params": request.optional_params, - })); - let (guardrail_request, _) = self - .guardrail_runner - .run_pre_call(&context, guardrail_request) - .await - .map_err(guardrail_error_to_core_error)?; - let (document, optional_params) = parse_ocr_pre_call_guardrail_request(guardrail_request)?; - let optional_params = match &request.config { - Ok(config) => config.map_ocr_params(&optional_params), - Err(_) => optional_params, - }; - Ok(PreparedOcrRequest { - document, - optional_params, - ..request - }) - } - - pub(crate) async fn prepare_provider_request( - &self, - request: PreparedOcrRequest, - ) -> Result { - let config = request.config?; - let env_lookup = |key: &str| std::env::var(key).ok(); - let upstream_headers = config.validate_environment( - string_headers(request.extra_headers)?, - request.api_key.as_deref(), - &env_lookup, - )?; - let url = config.complete_url( - request.api_base.as_deref(), - &request.model, - &request.optional_params, - &env_lookup, - )?; - let model = request.model.clone(); - let custom_llm_provider = request.custom_llm_provider.clone(); - let is_reducto = custom_llm_provider == "reducto"; - let document = if is_reducto { - let guarded_document = self - .run_during_call_guardrails(&model, &custom_llm_provider, &url, request.document) - .await?; - upload_reducto_document( - &guarded_document, - request.api_base.as_deref(), - request.timeout, - &upstream_headers, - ) - .await? - } else if config.requires_data_uri_document() { - convert_document_url_to_data_uri(request.document).await? - } else { - request.document - }; - let optional_params = request.optional_params; - let body = config - .transform_ocr_request(&request.model, document, optional_params.clone())? - .data; - let body = if is_reducto { - body - } else { - self.run_during_call_guardrails(&model, &custom_llm_provider, &url, body) - .await? - }; - Ok(ProviderOcrRequest { - model, - config, - url, - body, - optional_params, - upstream_headers, - timeout: request.timeout, - }) - } - - async fn run_during_call_guardrails( - &self, - model: &str, - custom_llm_provider: &str, - url: &str, - body: Value, - ) -> Result { - if self.guardrail_runner.is_empty() { - return Ok(body); - } - - let context = guardrail_context(&self.request_metadata); - let guardrail_request = GuardrailRequest::new(json!({ - "model": model, - "custom_llm_provider": custom_llm_provider, - "url": url, - "body": body, - })); - let (guardrail_request, _) = self - .guardrail_runner - .run_during_call(&context, guardrail_request) - .await - .map_err(guardrail_error_to_core_error)?; - parse_ocr_during_call_guardrail_request(guardrail_request) - } - - fn standard_logging_payload( - &self, - context: &CallLifecycleContext, - timing: &CallLifecycleTiming, - ) -> StandardLoggingPayload { - StandardLoggingPayload { - id: context.litellm_call_id.clone(), - litellm_call_id: context.litellm_call_id.clone(), - call_type: context.call_type.clone(), - model: context.model.clone(), - custom_llm_provider: context.custom_llm_provider.clone(), - response_cost: 0.0, - prompt_tokens: 0, - completion_tokens: 0, - total_tokens: 0, - start_time: timing.start_time, - end_time: timing.end_time, - stream: false, - metadata: StandardLoggingMetadata { - user_api_key_hash: self.request_metadata.user_api_key_hash.clone(), - user_api_key_user_id: self.request_metadata.user_api_key_user_id.clone(), - user_api_key_team_id: self.request_metadata.user_api_key_team_id.clone(), - ..Default::default() - }, - messages: None, - } - } -} - -async fn upload_reducto_document( - document: &Value, - api_base: Option<&str>, - timeout: Option, - upstream_headers: &[(String, String)], -) -> Result { - let source = extract_document_source(document)?; - let Some(authorization) = upstream_headers - .iter() - .find(|(name, _)| name.eq_ignore_ascii_case("authorization")) - .map(|(_, value)| value.as_str()) - else { - return Err(Error::Auth( - "Reducto upload requires an Authorization header".to_string(), - )); - }; - let Some(upload) = build_upload_request(source, authorization, api_base) else { - return Ok(document.clone()); - }; - let part = reqwest::multipart::Part::bytes(upload.bytes) - .file_name(upload.file_name) - .mime_str(&upload.mime_type) - .map_err(|error| Error::InvalidRequest(error.to_string()))?; - let form = reqwest::multipart::Form::new().part("file", part); - let mut request_builder = http_client().post(upload.url).multipart(form); - for (name, value) in upstream_headers { - if !name.eq_ignore_ascii_case("content-type") - && !name.eq_ignore_ascii_case("content-length") - { - request_builder = request_builder.header(name, value); - } - } - if let Some(timeout) = timeout { - request_builder = request_builder.timeout(timeout); - } - let response = request_builder - .send() - .await - .map_err(|error| Error::Network(error.to_string()))?; - let status = response.status(); - let body = response - .text() - .await - .map_err(|error| Error::Network(error.to_string()))?; - if !status.is_success() { - return Err(Error::Http { - status: status.as_u16(), - body: truncate_error_body(&body), - }); - } - let response_json: Value = serde_json::from_str(&body).map_err(|error| { - Error::InvalidResponse(format!("invalid Reducto upload response JSON: {error}")) - })?; - let file_id = extract_upload_file_id(&response_json)?; - Ok(json!({"type": "document_url", "document_url": file_id})) -} - -impl CallLifecycleHooks for OcrLifecycleHooks { - type PreCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; - type DuringCallFuture<'a> = OcrFuture<'a, PreparedOcrRequest>; - type SuccessFuture<'a> = OcrLogFuture<'a>; - type FailureFuture<'a> = OcrLogFuture<'a>; - - fn async_pre_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: PreparedOcrRequest, - ) -> Self::PreCallFuture<'a> { - Box::pin(async move { self.run_pre_call_guardrails(request).await }) - } - - fn async_during_call_hook<'a>( - &'a self, - _context: &'a CallLifecycleContext, - request: PreparedOcrRequest, - ) -> Self::DuringCallFuture<'a> { - Box::pin(async move { Ok(request) }) - } - - #[tracing::instrument( - name = "success_callback", - target = "litellm::function_trace", - level = "trace", - skip_all - )] - fn async_log_success_event<'a>( - &'a self, - context: &'a CallLifecycleContext, - response: &'a Value, - timing: &'a CallLifecycleTiming, - ) -> Self::SuccessFuture<'a> { - Box::pin(async move { - if self.logger_runner.is_empty() { - return; - } - let response_obj = CallbackValue::new("ocr", response.clone()); - self.logger_runner - .async_log_success_event( - &ModelCallDetails::from_standard_logging_payload( - self.standard_logging_payload(context, timing), - ), - &response_obj, - CallbackTiming::new(timing.start_time, timing.end_time), - ) - .await; - }) - } - - #[tracing::instrument( - name = "failure_callback", - target = "litellm::function_trace", - level = "trace", - skip_all - )] - fn async_log_failure_event<'a>( - &'a self, - context: &'a CallLifecycleContext, - error: &'a Error, - timing: &'a CallLifecycleTiming, - ) -> Self::FailureFuture<'a> { - Box::pin(async move { - if self.logger_runner.is_empty() { - return; - } - let logging_error = LoggingError { - message: error.to_string(), - kind: core_error_kind(error).to_string(), - }; - let response_obj = CallbackValue::new( - "error", - json!({ - "message": logging_error.message, - "kind": logging_error.kind, - }), - ); - self.logger_runner - .async_log_failure_event( - &ModelCallDetails::from_standard_logging_payload( - self.standard_logging_payload(context, timing), - ) - .with_failure_error(logging_error), - Some(&response_obj), - CallbackTiming::new(timing.start_time, timing.end_time), - ) - .await; - }) - } -} - -fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext { - GuardrailContext { - call_type: CallType::Ocr, - selected_guardrails: Vec::new(), - metadata: std::collections::HashMap::new(), - user_api_key_hash: metadata.user_api_key_hash.clone(), - user_api_key_user_id: metadata.user_api_key_user_id.clone(), - user_api_key_team_id: metadata.user_api_key_team_id.clone(), - trace_parent: None, - } -} - -fn parse_ocr_pre_call_guardrail_request( - request: GuardrailRequest, -) -> Result<(Value, Map), Error> { - let Value::Object(mut data) = request.data else { - return Err(Error::InvalidRequest( - "OCR pre_call guardrail must return an object".to_string(), - )); - }; - let document = data.remove("document").ok_or_else(|| { - Error::InvalidRequest("OCR pre_call guardrail removed document".to_string()) - })?; - let optional_params = match data.remove("optional_params") { - Some(Value::Object(params)) => params, - Some(_) => { - return Err(Error::InvalidRequest( - "OCR pre_call guardrail optional_params must be an object".to_string(), - )); - } - None => Map::new(), - }; - Ok((document, optional_params)) -} - -fn parse_ocr_during_call_guardrail_request(request: GuardrailRequest) -> Result { - let Value::Object(mut data) = request.data else { - return Err(Error::InvalidRequest( - "OCR during_call guardrail must return an object".to_string(), - )); - }; - data.remove("body") - .ok_or_else(|| Error::InvalidRequest("OCR during_call guardrail removed body".to_string())) -} - -fn guardrail_error_to_core_error(error: GuardrailError) -> Error { - Error::InvalidRequest(format!("{}: {}", error.kind, error.message)) -} - -fn core_error_kind(error: &Error) -> &'static str { - match error { - Error::Auth(_) | Error::MissingApiKey { .. } => "AuthError", - Error::InvalidProvider(_) => "InvalidProvider", - Error::InvalidRequest(_) => "InvalidRequest", - Error::InvalidType { .. } => "InvalidType", - Error::MissingField(_) => "MissingField", - Error::Http { .. } => "HttpError", - Error::InvalidResponse(_) => "InvalidResponse", - Error::Network(_) => "NetworkError", - Error::Connect(_) => "ConnectError", - Error::Routing(_) => "RoutingError", - Error::Unsupported(_) => "UnsupportedRequest", - } -} diff --git a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs index 2acdd232c80..fb63a02f7ad 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs @@ -1,174 +1,127 @@ use litellm_core::Error; -use litellm_core::call_lifecycle::CallLifecycle; +use litellm_core::ocr::{ + OcrClient, + wire::{OcrWireRequest, decode_request}, +}; use serde_json::Value; -mod common_utils; -mod handler; -mod hooks; -mod prepare; mod types; pub use types::OcrRequest; -use handler::execute_ocr_provider_call; -use prepare::{PreparedOcrCall, prepare_ocr_call}; - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] pub async fn ocr(request: OcrRequest<'_>) -> Result { - let PreparedOcrCall { request, hooks } = prepare_ocr_call(request); - CallLifecycle::default() - .run_request(request, &hooks, |request| { - execute_ocr_provider_call(request, &hooks) - }) + core_ocr(request).await +} + +async fn core_ocr(request: OcrRequest<'_>) -> Result { + validate_host_hooks(&request)?; + let client = OcrClient::new(crate::client::http_client().clone())?; + let core_request = decode_request(OcrWireRequest { + model: request.model.to_string(), + document: request.document, + api_key: request.api_key.map(str::to_string), + api_base: request.api_base.map(str::to_string), + custom_llm_provider: request.custom_llm_provider.map(str::to_string), + extra_headers: request.extra_headers, + optional_params: request.optional_params, + input_sources: Default::default(), + timeout_seconds: request.timeout.map(|timeout| timeout.as_secs_f64()), + })?; + client + .perform(core_request) .await + .map(|response| response.into_json()) +} + +fn validate_host_hooks(request: &OcrRequest<'_>) -> Result<(), Error> { + if !request.guardrails.is_empty() { + return Err(Error::Unsupported( + "OCR host guardrails are not wired to the core path", + )); + } + if !request.callbacks.is_empty() { + return Err(Error::Unsupported( + "OCR host callbacks are not wired to the core path", + )); + } + Ok(()) } #[cfg(test)] mod tests { + use std::sync::Arc; + + use litellm_core::ocr::wire::is_supported_request; use serde_json::{Map, json}; - use tokio::io::{AsyncReadExt, AsyncWriteExt}; - use tokio::net::{TcpListener, TcpStream}; - use super::{OcrRequest, ocr}; - use crate::integrations::types::RequestMetadata; + use super::{OcrRequest, validate_host_hooks}; + use crate::integrations::custom_guardrail::{CustomGuardrail, GuardrailEventHook}; + use crate::integrations::custom_logger::CustomLogger; - async fn read_http_request(socket: &mut TcpStream) -> String { - let mut request = Vec::new(); - let mut buffer = [0_u8; 1024]; - let header_end = loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break request.len(); - } - request.extend_from_slice(&buffer[..n]); - if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") { - break position + 4; - } - }; - let headers = String::from_utf8_lossy(&request[..header_end]); - let content_length = headers - .lines() - .find_map(|line| { - let (name, value) = line.split_once(':')?; - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().ok()) - .flatten() - }) - .unwrap_or(0); - while request.len().saturating_sub(header_end) < content_length { - let n = socket.read(&mut buffer).await.expect("reads body"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); + struct TestGuardrail; + + impl CustomGuardrail for TestGuardrail { + fn guardrail_name(&self) -> &str { + "test" + } + + fn supported_event_hooks(&self) -> &[GuardrailEventHook] { + &[] } - String::from_utf8(request).expect("request is utf8") } - fn base_ocr_request(model: &str) -> OcrRequest<'_> { + struct TestLogger; + + impl CustomLogger for TestLogger {} + + fn request() -> OcrRequest<'static> { OcrRequest { - model, - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), + model: "model", + document: json!({"type":"image_url","image_url":"data:image/png;base64,YQ=="}), + api_key: None, api_base: None, - custom_llm_provider: None, + custom_llm_provider: Some("mistral"), extra_headers: None, optional_params: Map::new(), timeout: None, callbacks: Vec::new(), guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), + request_metadata: Default::default(), litellm_call_id: None, } } - #[tokio::test] - async fn reducto_file_upload_then_parse_maps_response() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let address = listener.local_addr().expect("listener has local address"); - let server = tokio::spawn(async move { - let (mut upload_socket, _) = listener.accept().await.expect("accepts upload request"); - let upload_request = read_http_request(&mut upload_socket).await; - let upload_body = r#"{"file_id":"reducto://uploaded.pdf"}"#; - let upload_response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - upload_body.len(), - upload_body - ); - upload_socket - .write_all(upload_response.as_bytes()) - .await - .expect("writes upload response"); + #[test] + fn core_activation_includes_migrated_providers() { + assert!(is_supported_request("model", Some("mistral"))); + assert!(is_supported_request("pixtral-12b", Some("azure_ai"))); + assert!(is_supported_request( + "doc-intelligence/prebuilt-layout", + Some("azure_ai") + )); + assert!(is_supported_request("parse-v3", Some("reducto"))); + assert!(is_supported_request("mistral-ocr", Some("vertex_ai"))); + assert!(is_supported_request("deepseek-ocr", Some("vertex_ai"))); + } - let (mut parse_socket, _) = listener.accept().await.expect("accepts parse request"); - let parse_request = read_http_request(&mut parse_socket).await; - let parse_body = r#"{"job_id":"job_123","usage":{"num_pages":3,"credits":3},"result":{"chunks":[{"content":"Page 1 block A","blocks":[{"content":"Page 1 block A","bbox":{"page":1},"kind":"text"}]},{"content":"Page 2 block A","blocks":[{"content":"Page 2 block A","bbox":{"page":2},"kind":"table"}]},{"content":"Page 1 block B","blocks":[{"content":"Page 1 block B","bbox":{"page":1},"kind":"text"}]},{"content":"Page 3 block A","blocks":[{"content":"Page 3 block A","bbox":{"page":3},"kind":"figure"}]}]}}"#; - let parse_response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - parse_body.len(), - parse_body - ); - parse_socket - .write_all(parse_response.as_bytes()) - .await - .expect("writes parse response"); - (upload_request, parse_request) - }); - let api_base = format!("http://{address}"); - let mut request = base_ocr_request("reducto/parse-v3"); - request.api_base = Some(&api_base); - request.api_key = None; - request.extra_headers = Some(Map::from_iter([ - ("Authorization".to_string(), json!("Bearer test-key")), - ("x-trace-id".to_string(), json!("trace-1")), - ])); - request.document = json!({ - "type": "document_url", - "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" - }); - request.optional_params = Map::from_iter([ - ( - "formatting".to_string(), - json!({"table_output_format": "html"}), - ), - ("retrieval".to_string(), json!({"chunk_mode": "section"})), - ("settings".to_string(), json!({"ocr_system": "standard"})), - ]); + #[test] + fn core_path_rejects_unwired_guardrails() { + let request = OcrRequest { + guardrails: vec![Arc::new(TestGuardrail)], + ..request() + }; + let error = validate_host_hooks(&request).unwrap_err(); + assert!(error.to_string().contains("guardrails are not wired")); + } - let response = ocr(request).await.expect("Reducto OCR succeeds"); - - assert_eq!(response["pages"].as_array().map(Vec::len), Some(3)); - assert_eq!( - response["pages"][0]["markdown"], - "Page 1 block A\n\nPage 1 block B" - ); - assert_eq!(response["pages"][1]["markdown"], "Page 2 block A"); - assert_eq!(response["pages"][2]["markdown"], "Page 3 block A"); - assert_eq!(response["usage_info"]["pages_processed"], 3); - assert_eq!(response["usage_info"]["credits"], 3); - assert_eq!(response["provider_native_response"]["job_id"], "job_123"); - let (upload_request, parse_request) = server.await.expect("server task completes"); - assert!( - upload_request - .to_ascii_lowercase() - .contains("authorization: bearer test-key") - ); - assert!(upload_request.contains("application/pdf")); - assert!(upload_request.contains("%PDF-1.4")); - assert!(upload_request.contains("x-trace-id: trace-1")); - assert!( - parse_request - .to_ascii_lowercase() - .contains("authorization: bearer test-key") - ); - assert!(parse_request.contains(r#""input":"reducto://uploaded.pdf""#)); - assert!(parse_request.contains(r#""table_output_format":"html""#)); - assert!(parse_request.contains(r#""chunk_mode":"section""#)); - assert!(parse_request.contains(r#""ocr_system":"standard""#)); + #[test] + fn core_path_rejects_unwired_callbacks() { + let request = OcrRequest { + callbacks: vec![Arc::new(TestLogger)], + ..request() + }; + let error = validate_host_hooks(&request).unwrap_err(); + assert!(error.to_string().contains("callbacks are not wired")); } } diff --git a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs deleted file mode 100644 index fa9ca1a193e..00000000000 --- a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs +++ /dev/null @@ -1,163 +0,0 @@ -use std::sync::atomic::{AtomicU64, Ordering}; -use std::time::{SystemTime, UNIX_EPOCH}; - -use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; -use serde_json::{Map, Value}; - -use super::common_utils::ocr_provider_config; -use super::hooks::OcrLifecycleHooks; -use super::types::{OcrRequest, PreparedOcrRequest}; -use crate::integrations::custom_guardrail::CustomGuardrailRunner; -use crate::integrations::custom_logger::CustomLoggerRunner; - -pub(crate) struct PreparedOcrCall { - pub(crate) request: PreparedOcrRequest, - pub(crate) hooks: OcrLifecycleHooks, -} - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub(crate) fn prepare_ocr_call(request: OcrRequest<'_>) -> PreparedOcrCall { - let call_id = request - .litellm_call_id - .map(str::to_string) - .unwrap_or_else(new_ocr_call_id); - let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) - .unwrap_or(CustomLlmProvider { - model: request.model, - custom_llm_provider: "mistral", - }); - let model = provider_info.model.to_string(); - let custom_llm_provider = provider_info.custom_llm_provider.to_string(); - let config = ocr_provider_config(&custom_llm_provider, &model) - .ok_or_else(|| litellm_core::Error::InvalidProvider(custom_llm_provider.clone())) - .and_then(|config| { - validate_request_format(config, &request.optional_params, &custom_llm_provider)?; - Ok(config) - }); - let optional_params = match &config { - Ok(config) => { - let supported = config.supported_ocr_params(); - let mut mapped = config.map_ocr_params( - &request - .optional_params - .iter() - .filter(|(name, _)| supported.contains(&name.as_str())) - .map(|(name, value)| (name.clone(), value.clone())) - .collect(), - ); - for name in [ - "vertex_project", - "vertex_ai_project", - "vertex_location", - "vertex_ai_location", - ] { - if let Some(value) = request.optional_params.get(name) { - mapped.insert(name.to_string(), value.clone()); - } - } - mapped - } - Err(_) => request.optional_params, - }; - - PreparedOcrCall { - request: PreparedOcrRequest { - config, - model, - custom_llm_provider, - litellm_call_id: call_id, - document: request.document, - api_key: request.api_key.map(str::to_string), - api_base: request.api_base.map(str::to_string), - extra_headers: request.extra_headers, - optional_params, - timeout: request.timeout, - }, - hooks: OcrLifecycleHooks::new( - CustomLoggerRunner::new(request.callbacks), - CustomGuardrailRunner::new(request.guardrails), - request.request_metadata, - ), - } -} - -fn validate_request_format( - config: &'static dyn litellm_core::ocr::transformation::OcrProviderConfig, - optional_params: &Map, - provider: &str, -) -> Result<(), litellm_core::Error> { - let Some(format) = optional_params.get("req_format") else { - return Ok(()); - }; - match format.as_str() { - Some("litellm") => Ok(()), - Some("native") if config.supported_ocr_params().contains(&"req_format") => Ok(()), - Some("native") => Err(litellm_core::Error::InvalidRequest(format!( - "`req_format=native` is not supported for provider {provider}" - ))), - _ => Err(litellm_core::Error::InvalidRequest(format!( - "Invalid `req_format`: {format}. Expected `litellm` or `native`" - ))), - } -} - -fn new_ocr_call_id() -> String { - static COUNTER: AtomicU64 = AtomicU64::new(1); - let sequence = COUNTER.fetch_add(1, Ordering::Relaxed); - let timestamp = SystemTime::now() - .duration_since(UNIX_EPOCH) - .map(|duration| duration.as_nanos()) - .unwrap_or(0); - format!("ocr-{timestamp}-{sequence}") -} - -#[cfg(test)] -mod tests { - use litellm_core::error::Error; - use serde_json::{Map, json}; - - use super::{OcrRequest, prepare_ocr_call}; - use crate::integrations::types::RequestMetadata; - - fn base_ocr_request(model: &str) -> OcrRequest<'_> { - OcrRequest { - model, - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: Map::new(), - timeout: None, - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - litellm_call_id: None, - } - } - - fn request_with_format(format: &str) -> OcrRequest<'_> { - let mut request = base_ocr_request("mistral/mistral-ocr-latest"); - request.optional_params = Map::from_iter([("req_format".to_string(), json!(format))]); - request - } - - #[test] - fn native_format_rejected_for_provider_without_support_as_bad_request() { - let prepared = prepare_ocr_call(request_with_format("native")); - assert!( - matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("not supported for provider")) - ); - } - - #[test] - fn unknown_format_rejected_for_provider_without_support_as_bad_request() { - let prepared = prepare_ocr_call(request_with_format("raw")); - assert!( - matches!(prepared.request.config, Err(Error::InvalidRequest(message)) if message.contains("Invalid `req_format`")) - ); - } -} diff --git a/litellm-rust/crates/ai-gateway/src/ocr/types.rs b/litellm-rust/crates/ai-gateway/src/ocr/types.rs index 75a8e61ddbf..e96d2df1adb 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/types.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/types.rs @@ -1,8 +1,6 @@ use std::sync::Arc; use std::time::Duration; -use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest}; -use litellm_core::ocr::transformation::OcrProviderConfig; use serde_json::{Map, Value}; use crate::integrations::custom_guardrail::CustomGuardrail; @@ -23,37 +21,3 @@ pub struct OcrRequest<'a> { pub request_metadata: RequestMetadata, pub litellm_call_id: Option<&'a str>, } - -pub(crate) struct PreparedOcrRequest { - pub(crate) config: Result<&'static dyn OcrProviderConfig, litellm_core::Error>, - pub(crate) model: String, - pub(crate) custom_llm_provider: String, - pub(crate) litellm_call_id: String, - pub(crate) document: Value, - pub(crate) api_key: Option, - pub(crate) api_base: Option, - pub(crate) extra_headers: Option>, - pub(crate) optional_params: Map, - pub(crate) timeout: Option, -} - -impl CallLifecycleRequest for PreparedOcrRequest { - fn lifecycle_context(&self) -> CallLifecycleContext { - CallLifecycleContext::new( - "ocr", - self.model.clone(), - self.custom_llm_provider.clone(), - self.litellm_call_id.clone(), - ) - } -} - -pub(crate) struct ProviderOcrRequest { - pub(crate) model: String, - pub(crate) config: &'static dyn OcrProviderConfig, - pub(crate) url: String, - pub(crate) body: Value, - pub(crate) optional_params: Map, - pub(crate) upstream_headers: Vec<(String, String)>, - pub(crate) timeout: Option, -} diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs index c22d05f5726..39465e28e84 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -105,7 +105,11 @@ impl IntoResponse for MessagesRouteError { StatusCode::NOT_FOUND, "no messages deployment is configured for this model".to_string(), ), - Error::Auth(_) => ( + Error::Auth(_) + | Error::MissingApiKey { .. } + | Error::MissingAzureAiCredentials + | Error::MissingAzureDocumentIntelligenceCredentials + | Error::MissingReductoApiKey => ( StatusCode::BAD_GATEWAY, "messages provider authentication failed".to_string(), ), @@ -114,8 +118,7 @@ impl IntoResponse for MessagesRouteError { | Error::Connect(_) | Error::InvalidResponse(_) | Error::InvalidType { .. } - | Error::MissingField(_) - | Error::MissingApiKey { .. } => ( + | Error::MissingField(_) => ( StatusCode::BAD_GATEWAY, "messages provider request failed".to_string(), ), diff --git a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs b/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs deleted file mode 100644 index 60e90ed2a7c..00000000000 --- a/litellm-rust/crates/ai-gateway/tests/ocr_lifecycle.rs +++ /dev/null @@ -1,641 +0,0 @@ -use std::sync::{Arc, Mutex}; -use std::time::Duration; - -use litellm_ai_gateway::integrations::custom_guardrail::{ - CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook, - GuardrailFuture, GuardrailRequest, -}; -use litellm_ai_gateway::integrations::custom_logger::{ - CallbackTiming, CallbackValue, CustomLogger, LogFuture, ModelCallDetails, -}; -use litellm_ai_gateway::integrations::types::RequestMetadata; -use litellm_ai_gateway::ocr::{OcrRequest, ocr}; -use litellm_core::error::Error; -#[cfg(feature = "trace-parity")] -use litellm_core::observability::FunctionTrace; -use serde_json::{Map, Value, json}; -use tokio::io::{AsyncReadExt, AsyncWriteExt}; -use tokio::net::{TcpListener, TcpStream}; -#[cfg(feature = "trace-parity")] -use tracing::instrument::WithSubscriber; - -async fn read_http_headers(socket: &mut TcpStream) -> String { - let mut request = Vec::new(); - let mut buffer = [0_u8; 1024]; - loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); - if request.windows(4).any(|window| window == b"\r\n\r\n") { - break; - } - } - String::from_utf8(request).expect("request is utf8") -} - -async fn read_http_request(socket: &mut TcpStream) -> String { - let mut request = Vec::new(); - let mut buffer = [0_u8; 1024]; - let header_end = loop { - let n = socket.read(&mut buffer).await.expect("reads request"); - if n == 0 { - break request.len(); - } - request.extend_from_slice(&buffer[..n]); - if let Some(position) = request.windows(4).position(|window| window == b"\r\n\r\n") { - break position + 4; - } - }; - let headers = String::from_utf8_lossy(&request[..header_end]); - let content_length = headers - .lines() - .find_map(|line| { - let (name, value) = line.split_once(':')?; - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().ok()) - .flatten() - }) - .unwrap_or(0); - while request.len().saturating_sub(header_end) < content_length { - let n = socket.read(&mut buffer).await.expect("reads body"); - if n == 0 { - break; - } - request.extend_from_slice(&buffer[..n]); - } - String::from_utf8(request).expect("request is utf8") -} - -#[derive(Clone, Debug, PartialEq)] -struct RecordedLogEvent { - hook: &'static str, - model: String, - call_type: String, - user_id: Option, - response_object: Option, - error_kind: Option, -} - -#[derive(Default)] -struct RecordingOcrLogger { - events: Mutex>, -} - -impl RecordingOcrLogger { - fn events(&self) -> Vec { - self.events.lock().unwrap().clone() - } -} - -impl CustomLogger for RecordingOcrLogger { - fn async_log_success_event<'a>( - &'a self, - model_call_details: &'a ModelCallDetails, - response_obj: &'a CallbackValue, - _timing: CallbackTiming, - ) -> LogFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push(RecordedLogEvent { - hook: "async_log_success_event", - model: model_call_details.model.clone(), - call_type: model_call_details.call_type.to_string(), - user_id: model_call_details.metadata.user_api_key_user_id.clone(), - response_object: Some(response_obj.object.clone()), - error_kind: None, - }); - Ok(()) - }) - } - - fn async_log_failure_event<'a>( - &'a self, - model_call_details: &'a ModelCallDetails, - response_obj: Option<&'a CallbackValue>, - _timing: CallbackTiming, - ) -> LogFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push(RecordedLogEvent { - hook: "async_log_failure_event", - model: model_call_details.model.clone(), - call_type: model_call_details.call_type.to_string(), - user_id: model_call_details.metadata.user_api_key_user_id.clone(), - response_object: response_obj.map(|value| value.object.clone()), - error_kind: model_call_details - .failure_error - .as_ref() - .map(|error| error.kind.clone()), - }); - Ok(()) - }) - } -} - -struct RecordingOcrGuardrail { - hooks: Vec, - events: Mutex>, - block_pre_call: bool, - block_during_call: bool, -} - -impl RecordingOcrGuardrail { - fn new(hooks: Vec) -> Self { - Self { - hooks, - events: Mutex::new(Vec::new()), - block_pre_call: false, - block_during_call: false, - } - } - - fn blocking_pre_call() -> Self { - Self { - hooks: vec![GuardrailEventHook::PreCall], - events: Mutex::new(Vec::new()), - block_pre_call: true, - block_during_call: false, - } - } - - fn blocking_during_call() -> Self { - Self { - hooks: vec![GuardrailEventHook::DuringCall], - events: Mutex::new(Vec::new()), - block_pre_call: false, - block_during_call: true, - } - } - - fn events(&self) -> Vec<&'static str> { - self.events.lock().unwrap().clone() - } -} - -impl CustomGuardrail for RecordingOcrGuardrail { - fn guardrail_name(&self) -> &str { - "recording-ocr-guardrail" - } - - fn supported_event_hooks(&self) -> &[GuardrailEventHook] { - &self.hooks - } - - fn async_pre_call_hook<'a>( - &'a self, - _context: &'a GuardrailContext, - mut request: GuardrailRequest, - ) -> GuardrailFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("async_pre_call_hook"); - if self.block_pre_call { - return Ok(GuardrailDecision::Block(GuardrailError::blocked( - "blocked before provider", - ))); - } - request.data["document"]["guarded_pre"] = json!(true); - Ok(GuardrailDecision::Mask(request)) - }) - } - - fn async_moderation_hook<'a>( - &'a self, - _context: &'a GuardrailContext, - mut request: GuardrailRequest, - ) -> GuardrailFuture<'a> { - Box::pin(async move { - self.events.lock().unwrap().push("async_moderation_hook"); - if self.block_during_call { - return Ok(GuardrailDecision::Block(GuardrailError::blocked( - "blocked before provider", - ))); - } - request.data["body"]["guarded_during"] = json!(true); - Ok(GuardrailDecision::Mask(request)) - }) - } -} - -fn base_ocr_request(model: &str) -> OcrRequest<'_> { - OcrRequest { - model, - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), - api_base: None, - custom_llm_provider: None, - extra_headers: None, - optional_params: Map::new(), - timeout: None, - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - litellm_call_id: None, - } -} - -#[tokio::test] -async fn reducto_during_call_guardrail_blocks_before_upload() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let address = listener.local_addr().expect("listener has local address"); - let api_base = format!("http://{address}"); - let guardrail = Arc::new(RecordingOcrGuardrail::blocking_during_call()); - let mut request = base_ocr_request("reducto/parse-v3"); - request.api_base = Some(&api_base); - request.document = json!({ - "type": "document_url", - "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" - }); - request.guardrails = vec![guardrail.clone()]; - - let error = ocr(request).await.expect_err("guardrail blocks upload"); - - assert!(matches!(error, Error::InvalidRequest(_))); - assert_eq!(guardrail.events(), vec!["async_moderation_hook"]); - let accepted = tokio::time::timeout(Duration::from_millis(100), listener.accept()).await; - assert!(accepted.is_err(), "upload socket should not be touched"); -} - -#[tokio::test] -async fn reducto_upload_error_body_is_truncated() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let address = listener.local_addr().expect("listener has local address"); - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts upload request"); - let _request = read_http_request(&mut socket).await; - let body = "x".repeat(300); - let response = format!( - "HTTP/1.1 500 Internal Server Error\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - body.len(), - body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes upload response"); - }); - let api_base = format!("http://{address}"); - let mut request = base_ocr_request("reducto/parse-v3"); - request.api_base = Some(&api_base); - request.document = json!({ - "type": "document_url", - "document_url": "data:application/pdf;base64,JVBERi0xLjQ=" - }); - - let error = ocr(request).await.expect_err("upload should fail"); - - assert!( - matches!(error, Error::Http { status: 500, body } if body.chars().count() < 300 && body.ends_with("... (truncated)")) - ); - server.await.expect("server task completes"); -} - -#[tokio::test] -async fn ocr_lifecycle_runs_pre_during_and_success_hooks() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts one request"); - let request = read_http_request(&mut socket).await; - let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#; - let response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - request - }); - - let logger = Arc::new(RecordingOcrLogger::default()); - let guardrail = Arc::new(RecordingOcrGuardrail::new(vec![ - GuardrailEventHook::PreCall, - GuardrailEventHook::DuringCall, - ])); - #[cfg(feature = "trace-parity")] - let trace = FunctionTrace::default(); - let api_base = format!("http://{addr}"); - let call = ocr(OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), - api_base: Some(&api_base), - custom_llm_provider: Some("mistral"), - extra_headers: None, - optional_params: Map::new(), - timeout: Some(Duration::from_secs(5)), - callbacks: vec![logger.clone()], - guardrails: vec![guardrail.clone()], - request_metadata: RequestMetadata { - user_api_key_user_id: Some("user-1".to_string()), - ..Default::default() - }, - litellm_call_id: Some("ocr-call-1"), - }); - #[cfg(feature = "trace-parity")] - let call = call.with_subscriber(trace.dispatcher()); - let response = call.await.expect("ocr request succeeds"); - - assert_eq!(response["pages"][0]["markdown"], "ok"); - assert_eq!( - guardrail.events(), - vec!["async_pre_call_hook", "async_moderation_hook"] - ); - assert_eq!( - logger.events(), - vec![RecordedLogEvent { - hook: "async_log_success_event", - model: "mistral-ocr-latest".to_string(), - call_type: "ocr".to_string(), - user_id: Some("user-1".to_string()), - response_object: Some("ocr".to_string()), - error_kind: None, - }] - ); - #[cfg(feature = "trace-parity")] - assert_eq!( - trace - .events() - .iter() - .filter(|event| event.function.ends_with("_callback")) - .map(|event| event.function) - .collect::>(), - vec!["success_callback"] - ); - - let request = server.await.expect("server task completes"); - assert!(request.contains(r#""guarded_pre":true"#), "{request}"); - assert!(request.contains(r#""guarded_during":true"#), "{request}"); -} - -#[tokio::test] -async fn ocr_lifecycle_runs_failure_hook_on_provider_error() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts one request"); - let _request = read_http_request(&mut socket).await; - let response_body = "provider failed"; - let response = format!( - "HTTP/1.1 500 Internal Server Error\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - }); - - let logger = Arc::new(RecordingOcrLogger::default()); - #[cfg(feature = "trace-parity")] - let trace = FunctionTrace::default(); - let api_base = format!("http://{addr}"); - let call = ocr(OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), - api_base: Some(&api_base), - custom_llm_provider: Some("mistral"), - extra_headers: None, - optional_params: Map::new(), - timeout: Some(Duration::from_secs(5)), - callbacks: vec![logger.clone()], - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - litellm_call_id: Some("ocr-call-2"), - }); - #[cfg(feature = "trace-parity")] - let call = call.with_subscriber(trace.dispatcher()); - let err = call.await.expect_err("provider error propagates"); - - assert!(matches!(err, Error::Http { status: 500, .. })); - server.await.expect("server task completes"); - assert_eq!( - logger.events(), - vec![RecordedLogEvent { - hook: "async_log_failure_event", - model: "mistral-ocr-latest".to_string(), - call_type: "ocr".to_string(), - user_id: None, - response_object: Some("error".to_string()), - error_kind: Some("HttpError".to_string()), - }] - ); - #[cfg(feature = "trace-parity")] - assert_eq!( - trace - .events() - .iter() - .filter(|event| event.function.ends_with("_callback")) - .map(|event| event.function) - .collect::>(), - vec!["failure_callback"] - ); -} - -#[tokio::test] -async fn ocr_lifecycle_pre_call_block_skips_provider_socket() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - let logger = Arc::new(RecordingOcrLogger::default()); - let guardrail = Arc::new(RecordingOcrGuardrail::blocking_pre_call()); - - let err = ocr(OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-test"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("mistral"), - extra_headers: None, - optional_params: Map::new(), - timeout: Some(Duration::from_millis(100)), - callbacks: vec![logger.clone()], - guardrails: vec![guardrail.clone()], - request_metadata: RequestMetadata::default(), - litellm_call_id: Some("ocr-call-3"), - }) - .await - .expect_err("guardrail blocks request"); - - assert!(matches!(err, Error::InvalidRequest(_))); - assert_eq!(guardrail.events(), vec!["async_pre_call_hook"]); - assert_eq!( - logger.events(), - vec![RecordedLogEvent { - hook: "async_log_failure_event", - model: "mistral-ocr-latest".to_string(), - call_type: "ocr".to_string(), - user_id: None, - response_object: Some("error".to_string()), - error_kind: Some("InvalidRequest".to_string()), - }] - ); - let accepted = tokio::time::timeout(Duration::from_millis(100), listener.accept()).await; - assert!(accepted.is_err(), "provider socket should not be touched"); -} - -#[tokio::test] -async fn ocr_does_not_duplicate_authorization_header_when_header_is_supplied() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - - let server = tokio::spawn(async move { - let (mut socket, _) = listener.accept().await.expect("accepts one request"); - let request = read_http_headers(&mut socket).await; - let response_body = r#"{"pages":[{"index":0,"markdown":"ok"}],"model":"mistral-ocr-latest","usage_info":{"pages_processed":1}}"#; - let response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - socket - .write_all(response.as_bytes()) - .await - .expect("writes response"); - request - }); - - let mut headers = Map::new(); - headers.insert( - "Authorization".to_string(), - Value::String("Bearer sk-from-python".to_string()), - ); - headers.insert( - "x-trace-id".to_string(), - Value::String("trace-1".to_string()), - ); - - let response = ocr(OcrRequest { - model: "mistral-ocr-latest", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("sk-for-rust-fallback"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("mistral"), - extra_headers: Some(headers), - optional_params: Map::new(), - timeout: Some(Duration::from_secs(5)), - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - litellm_call_id: None, - }) - .await - .expect("ocr request succeeds"); - - assert_eq!(response["pages"][0]["markdown"], "ok"); - - let request = server.await.expect("server task completes"); - let authorization_count = request - .lines() - .filter(|line| line.to_ascii_lowercase().starts_with("authorization:")) - .count(); - assert_eq!(authorization_count, 1, "{request}"); - assert!( - request.contains("authorization: Bearer sk-from-python") - || request.contains("Authorization: Bearer sk-from-python"), - "{request}" - ); -} - -#[tokio::test] -async fn document_intelligence_poll_uses_resolved_subscription_key() { - let listener = TcpListener::bind("127.0.0.1:0") - .await - .expect("test listener binds"); - let addr = listener.local_addr().expect("listener has local addr"); - let operation_url = format!("http://{addr}/operations/1"); - - let server = tokio::spawn(async move { - let (mut post_socket, _) = listener.accept().await.expect("accepts post request"); - let post_request = read_http_headers(&mut post_socket).await; - let post_response = format!( - "HTTP/1.1 202 Accepted\r\noperation-location: {operation_url}\r\ncontent-length: 0\r\nconnection: close\r\n\r\n" - ); - post_socket - .write_all(post_response.as_bytes()) - .await - .expect("writes post response"); - - let (mut poll_socket, _) = listener.accept().await.expect("accepts poll request"); - let poll_request = read_http_headers(&mut poll_socket).await; - let response_body = r#"{"status":"succeeded","analyzeResult":{"pages":[{"pageNumber":1,"lines":[{"content":"ok"}]}]}}"#; - let poll_response = format!( - "HTTP/1.1 200 OK\r\ncontent-type: application/json\r\ncontent-length: {}\r\nconnection: close\r\n\r\n{}", - response_body.len(), - response_body - ); - poll_socket - .write_all(poll_response.as_bytes()) - .await - .expect("writes poll response"); - (post_request, poll_request) - }); - - let response = ocr(OcrRequest { - model: "doc-intelligence/prebuilt-read", - document: json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }), - api_key: Some("di-key"), - api_base: Some(&format!("http://{addr}")), - custom_llm_provider: Some("azure_ai"), - extra_headers: None, - optional_params: Map::new(), - timeout: Some(Duration::from_secs(5)), - callbacks: Vec::new(), - guardrails: Vec::new(), - request_metadata: RequestMetadata::default(), - litellm_call_id: None, - }) - .await - .expect("document intelligence request succeeds"); - - assert_eq!(response["pages"][0]["markdown"], "ok"); - - let (post_request, poll_request) = server.await.expect("server task completes"); - assert!( - post_request - .to_ascii_lowercase() - .contains("ocp-apim-subscription-key: di-key"), - "{post_request}" - ); - assert!( - poll_request - .to_ascii_lowercase() - .contains("ocp-apim-subscription-key: di-key"), - "{poll_request}" - ); -} diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index b4ca88cf16f..a2433435e34 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -12,17 +12,25 @@ path = "tests/workspace_crate_allowlist.rs" [dependencies] base64.workspace = true +azure_core.workspace = true +azure_identity.workspace = true +data-url = "0.3.2" +gcp_auth.workspace = true +moka.workspace = true rand.workspace = true reqwest.workspace = true serde.workspace = true serde_json.workspace = true serde_path_to_error = "0.1" +strum.workspace = true +subtle.workspace = true tokio.workspace = true thiserror.workspace = true tracing.workspace = true tracing-subscriber = { workspace = true, optional = true } sha2.workspace = true url.workspace = true +veil.workspace = true aws-config = { version = "1.9.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } aws-credential-types = { version = "1.3.0", features = ["hardcoded-credentials"], optional = true } aws-sdk-sts = { version = "1.108.0", default-features = false, features = ["rustls", "rt-tokio"], optional = true } @@ -44,5 +52,4 @@ observability = ["dep:tracing-subscriber"] [dev-dependencies] rstest.workspace = true -tokio = { workspace = true, features = ["macros", "rt-multi-thread"] } tracing-subscriber.workspace = true diff --git a/litellm-rust/crates/core/src/auth/credential.rs b/litellm-rust/crates/core/src/auth/credential.rs new file mode 100644 index 00000000000..b5235b6780c --- /dev/null +++ b/litellm-rust/crates/core/src/auth/credential.rs @@ -0,0 +1,168 @@ +use std::future::Future; +use std::path::PathBuf; +use std::pin::Pin; +use std::sync::Arc; + +use veil::Redact; + +use crate::AuthError; + +use super::{ResolvedCredential, SecretValue, TokenProviderHandle}; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum CredentialFileRef { + Path(PathBuf), + EnvironmentVariable(String), +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum CredentialRef { + Explicit(SecretValue), + Env(String), + File(CredentialFileRef), + Request(String), + Host(String), + None, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum CredentialLookup { + Found(SecretValue), + Missing, + Declined, +} + +pub type CredentialLookupFuture<'a> = + Pin> + Send + 'a>>; + +pub trait CredentialResolver: std::fmt::Debug + Send + Sync { + fn resolve<'a>(&'a self, reference: &'a CredentialRef) -> CredentialLookupFuture<'a>; +} + +#[derive(Clone, Redact)] +pub struct CredentialResolverHandle(#[redact(with = "[REDACTED]")] Arc); + +impl CredentialResolverHandle { + pub fn new(resolver: Arc) -> Self { + Self(resolver) + } + + pub async fn resolve(&self, reference: &CredentialRef) -> Result { + self.0.resolve(reference).await + } +} + +#[derive(Clone, Debug)] +pub enum CredentialPlan { + Static(CredentialRef), + Caller(TokenProviderHandle), + None, +} + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum CredentialPlanResolution { + Resolved(ResolvedCredential), + Unavailable, +} + +impl CredentialPlan { + pub async fn resolve( + &self, + resolver: &CredentialResolverHandle, + ) -> Result { + match self { + Self::Static(CredentialRef::Explicit(secret)) => Ok( + CredentialPlanResolution::Resolved(ResolvedCredential::Static(secret.clone())), + ), + Self::Static(CredentialRef::None) | Self::None => { + Ok(CredentialPlanResolution::Unavailable) + } + Self::Static(reference) => match resolver.resolve(reference).await? { + CredentialLookup::Found(secret) => Ok(CredentialPlanResolution::Resolved( + ResolvedCredential::Static(secret), + )), + CredentialLookup::Missing | CredentialLookup::Declined => { + Ok(CredentialPlanResolution::Unavailable) + } + }, + Self::Caller(caller) => { + let credential = caller.acquire().await?; + if credential.secret().expose().is_empty() { + return Err(AuthError::EmptyCallerCredential); + } + Ok(CredentialPlanResolution::Resolved(credential)) + } + } + } +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use super::{ + CredentialLookup, CredentialLookupFuture, CredentialPlan, CredentialPlanResolution, + CredentialRef, CredentialResolver, CredentialResolverHandle, + }; + use crate::AuthError; + use crate::auth::SecretValue; + + #[derive(Debug)] + struct HostResolver; + + impl CredentialResolver for HostResolver { + fn resolve<'a>(&'a self, reference: &'a CredentialRef) -> CredentialLookupFuture<'a> { + Box::pin(async move { + Ok(match reference { + CredentialRef::Host(name) if name == "rotating-token" => { + CredentialLookup::Found(SecretValue::new("resolved")) + } + _ => CredentialLookup::Declined, + }) + }) + } + } + + #[tokio::test] + async fn static_host_reference_resolves_at_acquisition_time() { + let resolver = CredentialResolverHandle::new(Arc::new(HostResolver)); + let plan = CredentialPlan::Static(CredentialRef::Host("rotating-token".to_string())); + + let resolved = plan.resolve(&resolver).await.unwrap(); + + assert!(matches!(resolved, CredentialPlanResolution::Resolved(_))); + } + + #[tokio::test] + async fn declined_reference_is_available_for_pre_acquisition_fallback() { + let resolver = CredentialResolverHandle::new(Arc::new(HostResolver)); + let plan = CredentialPlan::Static(CredentialRef::Request("api-key".to_string())); + + assert_eq!( + plan.resolve(&resolver).await.unwrap(), + CredentialPlanResolution::Unavailable + ); + } + + #[derive(Debug)] + struct FailingResolver; + + impl CredentialResolver for FailingResolver { + fn resolve<'a>(&'a self, _reference: &'a CredentialRef) -> CredentialLookupFuture<'a> { + Box::pin(async { Err(AuthError::UnresolvedOidcReference) }) + } + } + + #[tokio::test] + async fn acquisition_failure_is_terminal() { + let resolver = CredentialResolverHandle::new(Arc::new(FailingResolver)); + let plan = CredentialPlan::Static(CredentialRef::Host("token".to_string())); + + let error = plan + .resolve(&resolver) + .await + .expect_err("acquisition errors cannot become fallback"); + + assert_eq!(error, AuthError::UnresolvedOidcReference); + } +} diff --git a/litellm-rust/crates/core/src/auth/error.rs b/litellm-rust/crates/core/src/auth/error.rs new file mode 100644 index 00000000000..e7027c0df10 --- /dev/null +++ b/litellm-rust/crates/core/src/auth/error.rs @@ -0,0 +1,128 @@ +use thiserror::Error; + +#[derive(Clone, Debug, Error, PartialEq, Eq)] +pub enum AuthError { + #[error("invalid authentication configuration: {0}")] + Configuration(#[from] AuthConfigurationError), + #[error("credential acquisition failed: {0}")] + AzureTokenAcquisition(String), + #[error("credential acquisition failed: Vertex AI credentials: {0}")] + VertexTokenAcquisition(String), + #[error("credential acquisition failed: {}", .0.iter().map(ToString::to_string).collect::>().join("; "))] + CredentialChain(Vec), + #[error("credential caller failed: credential caller returned an empty credential")] + EmptyCallerCredential, + #[error("credential caller failed: Azure AD token provider returned an empty token")] + EmptyAzureToken, + #[error("credential acquisition failed: Azure OIDC reference did not resolve to a value")] + UnresolvedOidcReference, + #[error( + "Missing {provider} API Key - A call is being made to {provider} but no key is set either in the environment variables or via params" + )] + MissingApiKey { provider: &'static str }, + #[error( + "Missing {provider} API Base - Set {environment_variable} environment variable or pass api_base parameter" + )] + MissingApiBase { + provider: &'static str, + environment_variable: &'static str, + }, + #[error("{0}")] + MissingCredential(#[from] MissingCredential), + #[error("{0}")] + Aws(#[from] AwsAuthError), + #[error("invalid authentication header")] + InvalidHeader, +} + +#[derive(Clone, Debug, Error, PartialEq, Eq)] +pub enum AuthConfigurationError { + #[error("credential header already exists")] + ExistingCredentialHeader, + #[error("credential plan is not allowed by the provider auth policy")] + DisallowedCredentialPlan, + #[error("credential cannot be empty")] + EmptyCredential, + #[error("invalid Azure credential selector")] + InvalidAzureSelector, + #[error("ClientSecretCredential requires tenant_id, client_id, and client_secret")] + MissingClientSecretFields, + #[error("WorkloadIdentityCredential requires tenant_id")] + MissingWorkloadTenant, + #[error("WorkloadIdentityCredential requires client_id")] + MissingWorkloadClient, + #[error("WorkloadIdentityCredential requires azure_federated_token_file")] + MissingWorkloadTokenFile, + #[error("credential reference requires a host credential resolver")] + MissingHostResolver, + #[error("caller credential plan requires provider-specific inputs")] + MissingCallerInputs, + #[error("credential header {0} already exists")] + DuplicateHeader(&'static str), + #[error("{0} must be a string or null")] + InvalidFieldType(String), + #[error("unsupported OIDC reference")] + UnsupportedOidcReference, + #[error("{0} cannot be empty")] + EmptyReference(String), + #[error("Azure credential initialization failed: {0}")] + AzureCredentialInitialization(String), + #[error("Azure authority must be an HTTPS origin without credentials, query, or fragment")] + InvalidAzureAuthority, + #[error("request-controlled Azure auth inputs cannot be combined with host credentials")] + MixedAzureCredentialSources, + #[error("request-controlled Azure credential references are not allowed")] + RequestAzureCredentialReference, + #[error("host credentials cannot be sent to a request-controlled Azure endpoint")] + RequestAzureCredentialDestination, + #[error("credentials cannot be sent to a request-controlled Vertex AI endpoint")] + RequestVertexCredentialDestination, + #[error( + "request-controlled Vertex credentials must use the canonical Google OAuth token endpoint" + )] + RequestVertexTokenEndpoint, +} + +#[derive(Clone, Debug, Error, PartialEq, Eq)] +pub enum MissingCredential { + #[error( + "Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY environment variable" + )] + AnthropicApiKey, + #[error("Missing Azure API Key - Set `api_key` or the AZURE_API_KEY environment variable")] + AzureApiKey, + #[error( + "Missing Azure API Base - Set `api_base` or the AZURE_API_BASE environment variable. Expected format: https://.services.ai.azure.com/anthropic" + )] + AzureApiBase, + #[error( + "Missing OpenAI API Key - a realtime call is being made but no key was passed via params or the OPENAI_API_KEY environment variable" + )] + OpenAiRealtimeApiKey, + #[error( + "Missing OpenAI API Key - a Responses WebSocket call is being made but no key was passed via params or the OPENAI_API_KEY environment variable" + )] + OpenAiResponsesApiKey, +} + +#[derive(Clone, Debug, Error, PartialEq, Eq)] +pub enum AwsAuthError { + #[error("AWS profile credentials failed: {0}")] + Profile(String), + #[error("AWS default credentials failed: {0}")] + DefaultChain(String), + #[error("AWS role credentials failed: {0}")] + AssumeRole(String), + #[error("AWS web identity credentials failed: {0}")] + WebIdentity(String), + #[error("AWS web identity expiration was invalid: {0}")] + WebIdentityExpiration(String), + #[error("AWS signing parameters failed: {0}")] + SigningParameters(String), + #[error("AWS signable request failed: {0}")] + SignableRequest(String), + #[error("AWS request signing failed: {0}")] + Signing(String), + #[error("AWS web identity response had no credentials")] + MissingWebIdentityCredentials, +} diff --git a/litellm-rust/crates/core/src/auth/http.rs b/litellm-rust/crates/core/src/auth/http.rs new file mode 100644 index 00000000000..83931311550 --- /dev/null +++ b/litellm-rust/crates/core/src/auth/http.rs @@ -0,0 +1,86 @@ +use crate::AuthError; +use crate::auth::error::AuthConfigurationError; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum CredentialPlacement { + Bearer, + Header(&'static str), +} + +impl CredentialPlacement { + pub fn header_name(self) -> &'static str { + match self { + Self::Bearer => "Authorization", + Self::Header(name) => name, + } + } +} + +pub(crate) fn apply_credential( + headers: Vec<(String, String)>, + credential: &str, + placement: CredentialPlacement, +) -> Result, AuthError> { + if credential.trim().is_empty() { + return Err(AuthError::Configuration( + AuthConfigurationError::EmptyCredential, + )); + } + if headers + .iter() + .any(|(name, _)| name.eq_ignore_ascii_case(placement.header_name())) + { + return Err(AuthError::Configuration( + AuthConfigurationError::DuplicateHeader(placement.header_name()), + )); + } + let value = match placement { + CredentialPlacement::Bearer => format!("Bearer {credential}"), + CredentialPlacement::Header(_) => credential.to_string(), + }; + Ok( + std::iter::once((placement.header_name().to_string(), value)) + .chain(headers) + .collect(), + ) +} + +/// How the upstream call is authenticated. API-key strategies are resolved in +/// `prepare`; SigV4 needs the serialized body, so the handler signs it. +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum RequestAuth { + Header { name: &'static str, value: String }, + Bearer { token: String }, + AwsSigV4 { region: String }, +} + +#[cfg(test)] +mod tests { + use super::{CredentialPlacement, apply_credential}; + + #[test] + fn bearer_uses_authorization_header() { + let headers = apply_credential(Vec::new(), "key", CredentialPlacement::Bearer) + .expect("credential applies"); + + assert_eq!( + headers, + vec![("Authorization".to_string(), "Bearer key".to_string())] + ); + } + + #[test] + fn named_header_rejects_existing_value() { + let error = apply_credential( + vec![( + "ocp-apim-subscription-key".to_string(), + "caller-key".to_string(), + )], + "configured-key", + CredentialPlacement::Header("Ocp-Apim-Subscription-Key"), + ) + .expect_err("provider policy must handle existing credentials"); + + assert!(error.to_string().contains("already exists")); + } +} diff --git a/litellm-rust/crates/core/src/auth/mod.rs b/litellm-rust/crates/core/src/auth/mod.rs new file mode 100644 index 00000000000..35d9c676f65 --- /dev/null +++ b/litellm-rust/crates/core/src/auth/mod.rs @@ -0,0 +1,56 @@ +mod credential; +pub mod error; +pub(crate) mod vertex; +pub use error::AuthError; +pub(crate) 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 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/core/src/auth/policy.rs b/litellm-rust/crates/core/src/auth/policy.rs new file mode 100644 index 00000000000..b796dedf0d8 --- /dev/null +++ b/litellm-rust/crates/core/src/auth/policy.rs @@ -0,0 +1,114 @@ +use crate::AuthError; +use crate::auth::error::AuthConfigurationError; + +use super::http::apply_credential; +use super::{CredentialPlacement, ResolvedCredential}; + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum CredentialPlanKind { + Static, + Entra, + Caller, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct CredentialRule { + pub kind: CredentialPlanKind, + pub placement: CredentialPlacement, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub enum ExistingHeaderBehavior { + Preserve, + Reject, +} + +#[derive(Clone, Copy, Debug, PartialEq, Eq)] +pub struct ProviderAuthPolicy { + pub rules: &'static [CredentialRule], + pub accepted_existing_headers: &'static [&'static str], + pub existing_header_behavior: ExistingHeaderBehavior, + pub scope: Option<&'static str>, + pub audience: Option<&'static str>, +} + +impl ProviderAuthPolicy { + pub fn has_existing_credential(&self, headers: &[(String, String)]) -> bool { + headers.iter().any(|(name, _)| { + self.accepted_existing_headers + .iter() + .any(|accepted| name.eq_ignore_ascii_case(accepted)) + }) + } + + pub fn apply( + &self, + headers: Vec<(String, String)>, + kind: CredentialPlanKind, + credential: &ResolvedCredential, + ) -> Result, AuthError> { + if self.has_existing_credential(&headers) { + return match self.existing_header_behavior { + ExistingHeaderBehavior::Preserve => Ok(headers), + ExistingHeaderBehavior::Reject => Err(AuthError::Configuration( + AuthConfigurationError::ExistingCredentialHeader, + )), + }; + } + let rule = + self.rules + .iter() + .find(|rule| rule.kind == kind) + .ok_or(AuthError::Configuration( + AuthConfigurationError::DisallowedCredentialPlan, + ))?; + apply_credential(headers, credential.secret().expose(), rule.placement) + } +} + +#[cfg(test)] +mod tests { + use super::{CredentialPlanKind, CredentialRule, ExistingHeaderBehavior, ProviderAuthPolicy}; + use crate::auth::{CredentialPlacement, ResolvedCredential, SecretValue}; + + const RULES: &[CredentialRule] = &[CredentialRule { + kind: CredentialPlanKind::Static, + placement: CredentialPlacement::Header("x-api-key"), + }]; + const POLICY: ProviderAuthPolicy = ProviderAuthPolicy { + rules: RULES, + accepted_existing_headers: &["x-api-key"], + existing_header_behavior: ExistingHeaderBehavior::Preserve, + scope: None, + audience: None, + }; + + #[test] + fn rules_define_allowed_plans_and_credential_placement() { + let headers = POLICY + .apply( + Vec::new(), + CredentialPlanKind::Static, + &ResolvedCredential::Static(SecretValue::new("secret")), + ) + .unwrap(); + + assert_eq!( + headers, + vec![("x-api-key".to_string(), "secret".to_string())] + ); + } + + #[test] + fn unsupported_plan_is_rejected() { + let error = POLICY + .apply( + Vec::new(), + CredentialPlanKind::Entra, + &ResolvedCredential::Static(SecretValue::new("secret")), + ) + .unwrap_err(); + + assert!(error.to_string().contains("not allowed")); + } +} diff --git a/litellm-rust/crates/core/src/auth/secret.rs b/litellm-rust/crates/core/src/auth/secret.rs new file mode 100644 index 00000000000..3ecb0a835ee --- /dev/null +++ b/litellm-rust/crates/core/src/auth/secret.rs @@ -0,0 +1,41 @@ +use veil::Redact; + +#[derive(Redact, Clone)] +pub struct SecretValue(#[redact(with = "[REDACTED]")] String); + +impl SecretValue { + pub fn new(value: impl Into) -> Self { + Self(value.into()) + } + + pub fn expose(&self) -> &str { + &self.0 + } +} + +impl PartialEq for SecretValue { + fn eq(&self, other: &Self) -> bool { + subtle::ConstantTimeEq::ct_eq(self.0.as_bytes(), other.0.as_bytes()).into() + } +} + +impl Eq for SecretValue {} + +#[cfg(test)] +mod tests { + use super::SecretValue; + + #[test] + fn debug_redacts_plaintext() { + let debug = format!("{:?}", SecretValue::new("credential-value")); + + assert!(!debug.contains("credential-value")); + assert!(debug.contains("REDACTED")); + } + + #[test] + fn equality_compares_plaintext_values() { + assert_eq!(SecretValue::new("same"), SecretValue::new("same")); + assert_ne!(SecretValue::new("same"), SecretValue::new("different")); + } +} diff --git a/litellm-rust/crates/core/src/auth/token.rs b/litellm-rust/crates/core/src/auth/token.rs new file mode 100644 index 00000000000..cfc6b8f0d6b --- /dev/null +++ b/litellm-rust/crates/core/src/auth/token.rs @@ -0,0 +1,47 @@ +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; +use std::time::SystemTime; + +use veil::Redact; + +use crate::AuthError; + +use super::secret::SecretValue; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum ResolvedCredential { + Static(SecretValue), + AccessToken { + token: SecretValue, + expires_on: Option, + }, +} + +impl ResolvedCredential { + pub fn secret(&self) -> &SecretValue { + match self { + Self::Static(secret) | Self::AccessToken { token: secret, .. } => secret, + } + } +} + +pub type TokenFuture<'a> = + Pin> + Send + 'a>>; + +pub trait TokenProvider: std::fmt::Debug + Send + Sync { + fn acquire(&self) -> TokenFuture<'_>; +} + +#[derive(Clone, Redact)] +pub struct TokenProviderHandle(#[redact(with = "[REDACTED]")] Arc); + +impl TokenProviderHandle { + pub fn new(caller: Arc) -> Self { + Self(caller) + } + + pub async fn acquire(&self) -> Result { + self.0.acquire().await + } +} diff --git a/litellm-rust/crates/core/src/auth/vertex.rs b/litellm-rust/crates/core/src/auth/vertex.rs new file mode 100644 index 00000000000..00a0a7ea7ee --- /dev/null +++ b/litellm-rust/crates/core/src/auth/vertex.rs @@ -0,0 +1,592 @@ +use std::collections::BTreeMap; +use std::future::Future; +use std::path::Path; +use std::pin::Pin; +use std::sync::Arc; + +use gcp_auth::{CustomServiceAccount, TokenProvider}; +use moka::future::Cache; +use serde_json::{Map, Value}; +use sha2::{Digest, Sha256}; + +use crate::auth::error::AuthConfigurationError; +use crate::auth::http::apply_credential; +use crate::auth::{AuthError, CredentialPlacement, InputSource, SecretValue, Sourced}; + +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"; +const VERTEX_AI_API_KEY_ENV: &str = "VERTEX_AI_API_KEY"; +const VERTEXAI_API_KEY_ENV: &str = "VERTEXAI_API_KEY"; +const VERTEXAI_CREDENTIALS_ENV: &str = "VERTEXAI_CREDENTIALS"; +const VERTEXAI_PROJECT_ENV: &str = "VERTEXAI_PROJECT"; +const VERTEXAI_LOCATION_ENV: &str = "VERTEXAI_LOCATION"; +const VERTEX_LOCATION_ENV: &str = "VERTEX_LOCATION"; + +#[derive(Clone, Debug, Default)] +pub(crate) struct VertexConfig { + credentials: Option>, + project_id: Option, + location: Option, +} + +impl VertexConfig { + pub(crate) fn from_sourced_optional_params( + params: &Map, + sources: &BTreeMap, + ) -> Result { + Ok(Self { + credentials: 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"])?, + }) + } + + pub(crate) fn project_id(&self) -> Option<&str> { + self.project_id.as_deref() + } + + pub(crate) fn location(&self) -> Option<&str> { + self.location.as_deref() + } +} + +pub(crate) struct VertexEnvironment { + pub headers: Vec<(String, String)>, + pub project_id: String, +} + +struct VertexAccessToken { + token: String, + project_id: String, +} + +pub(crate) fn get_vertex_ai_project( + config: &VertexConfig, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option { + config + .project_id() + .map(str::to_string) + .or_else(|| non_empty_env(env_lookup, VERTEXAI_PROJECT_ENV)) +} + +pub(crate) fn get_vertex_ai_location( + config: &VertexConfig, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option { + config + .location() + .map(str::to_string) + .or_else(|| non_empty_env(env_lookup, VERTEXAI_LOCATION_ENV)) + .or_else(|| non_empty_env(env_lookup, VERTEX_LOCATION_ENV)) +} + +#[derive(Clone)] +pub(crate) struct VertexAuth { + providers: Cache>, + loader: Arc, +} + +impl Default for VertexAuth { + fn default() -> Self { + Self::new(Arc::new(GcpProviderLoader)) + } +} + +impl VertexAuth { + fn new(loader: Arc) -> Self { + Self { + providers: Cache::builder().max_capacity(64).build(), + loader, + } + } + + #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] + pub(crate) async fn validate_environment( + &self, + headers: Vec<(String, String)>, + api_key: Option<&str>, + config: &VertexConfig, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result { + let has_authorization = headers + .iter() + .any(|(name, _)| name.eq_ignore_ascii_case("Authorization")); + let static_token = api_key + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .or_else(|| non_empty_env(env_lookup, VERTEX_AI_API_KEY_ENV)) + .or_else(|| non_empty_env(env_lookup, VERTEXAI_API_KEY_ENV)); + let project_id = get_vertex_ai_project(config, env_lookup); + + if !has_authorization && static_token.is_none() { + let access = self.get_access_token(config, env_lookup).await?; + return Ok(VertexEnvironment { + headers: apply_credential(headers, &access.token, CredentialPlacement::Bearer)?, + project_id: project_id.unwrap_or(access.project_id), + }); + } + + let project_id = match project_id { + Some(project_id) => project_id, + None => { + self.load_provider(config, env_lookup) + .await? + .project_id() + .await? + } + }; + let headers = if has_authorization { + headers + } else { + apply_credential( + headers, + static_token.as_deref().expect("static token was checked"), + CredentialPlacement::Bearer, + )? + }; + Ok(VertexEnvironment { + headers, + project_id, + }) + } + + async fn get_access_token( + &self, + config: &VertexConfig, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result { + let provider = self.load_provider(config, env_lookup).await?; + let (token, project_id) = tokio::try_join!(provider.token(), provider.project_id())?; + Ok(VertexAccessToken { token, project_id }) + } + + async fn load_provider( + &self, + config: &VertexConfig, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result, AuthError> { + let source = credential_source(config, env_lookup); + let key = source.cache_key(); + self.providers + .try_get_with(key, self.loader.load(source)) + .await + .map_err(|error| (*error).clone()) + } +} + +trait VertexTokenSource: Send + Sync { + fn project_id(&self) -> VertexAuthFuture<'_, String>; + fn token(&self) -> VertexAuthFuture<'_, String>; +} + +trait VertexProviderLoader: Send + Sync { + fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc>; +} + +type VertexAuthFuture<'a, T> = Pin> + Send + 'a>>; + +struct GcpTokenSource(Arc); + +impl VertexTokenSource for GcpTokenSource { + fn project_id(&self) -> VertexAuthFuture<'_, String> { + Box::pin(async move { + self.0 + .project_id() + .await + .map(|project| project.to_string()) + .map_err(auth_acquisition_error) + }) + } + + fn token(&self) -> VertexAuthFuture<'_, String> { + Box::pin(async move { + self.0 + .token(&[CLOUD_PLATFORM_SCOPE]) + .await + .map(|token| token.as_str().to_string()) + .map_err(auth_acquisition_error) + }) + } +} + +struct GcpProviderLoader; + +impl VertexProviderLoader for GcpProviderLoader { + fn load(&self, source: CredentialSource) -> VertexAuthFuture<'_, Arc> { + Box::pin(async move { + let provider: Arc = match source { + CredentialSource::Inline(configured) => Arc::new( + CustomServiceAccount::from_json(validate_request_credentials( + configured.expose(), + )?) + .map_err(auth_acquisition_error)?, + ), + CredentialSource::Trusted(configured) => { + let configured = configured.expose(); + let service_account = if Path::new(configured).is_file() { + CustomServiceAccount::from_file(configured) + } else { + CustomServiceAccount::from_json(configured) + } + .map_err(auth_acquisition_error)?; + Arc::new(service_account) + } + CredentialSource::ApplicationCredentials(path) => { + Arc::new(CustomServiceAccount::from_file(path).map_err(auth_acquisition_error)?) + } + CredentialSource::Adc => { + gcp_auth::provider().await.map_err(auth_acquisition_error)? + } + }; + Ok(Arc::new(GcpTokenSource(provider)) as Arc) + }) + } +} + +fn validate_request_credentials(configured: &str) -> Result<&str, AuthError> { + let token_uri = serde_json::from_str::(configured) + .ok() + .and_then(|credentials| { + credentials + .get("token_uri") + .and_then(Value::as_str) + .map(str::to_string) + }); + if token_uri.as_deref() != Some(GOOGLE_OAUTH_TOKEN_ENDPOINT) { + return Err(AuthConfigurationError::RequestVertexTokenEndpoint.into()); + } + Ok(configured) +} + +#[derive(Clone, Debug)] +enum CredentialSource { + Inline(SecretValue), + Trusted(SecretValue), + ApplicationCredentials(String), + Adc, +} + +impl CredentialSource { + fn cache_key(&self) -> CredentialCacheKey { + match self { + Self::Inline(configured) => { + CredentialCacheKey::Inline(Sha256::digest(configured.expose()).into()) + } + Self::Trusted(configured) => { + CredentialCacheKey::Trusted(Sha256::digest(configured.expose()).into()) + } + Self::ApplicationCredentials(path) => { + CredentialCacheKey::ApplicationCredentials(path.clone()) + } + Self::Adc => CredentialCacheKey::Adc, + } + } +} + +#[derive(Clone, Debug, Hash, PartialEq, Eq)] +enum CredentialCacheKey { + Inline([u8; 32]), + Trusted([u8; 32]), + ApplicationCredentials(String), + Adc, +} + +fn credential_source( + config: &VertexConfig, + env_lookup: &dyn Fn(&str) -> Option, +) -> CredentialSource { + if let Some(configured) = config.credentials.clone() { + return match configured.source() { + InputSource::Request => CredentialSource::Inline(configured.into_value()), + InputSource::Deployment | InputSource::Environment => { + CredentialSource::Trusted(configured.into_value()) + } + }; + } + if let Some(configured) = non_empty_env(env_lookup, VERTEXAI_CREDENTIALS_ENV) { + return CredentialSource::Trusted(SecretValue::new(configured)); + } + non_empty_env(env_lookup, GOOGLE_APPLICATION_CREDENTIALS_ENV) + .map(CredentialSource::ApplicationCredentials) + .unwrap_or(CredentialSource::Adc) +} + +fn optional_credentials( + params: &Map, + sources: &BTreeMap, + names: &[&str], +) -> Result>, AuthError> { + for name in names { + let source = source_for(sources, name); + match params.get(*name) { + None | Some(Value::Null) => continue, + Some(Value::String(value)) if value.trim().is_empty() => continue, + Some(Value::String(value)) => { + return Ok(Some(Sourced::new(SecretValue::new(value), source))); + } + Some(Value::Object(value)) if value.is_empty() => continue, + Some(Value::Object(value)) => { + return serde_json::to_string(value) + .map(SecretValue::new) + .map(|value| Sourced::new(value, source)) + .map(Some) + .map_err(|error| { + AuthError::Configuration(AuthConfigurationError::InvalidFieldType(format!( + "{}: {error}", + names[0] + ))) + }); + } + Some(_) => { + return Err(AuthError::Configuration( + AuthConfigurationError::InvalidFieldType(names[0].to_string()), + )); + } + } + } + Ok(None) +} + +fn source_for(sources: &BTreeMap, name: &str) -> InputSource { + sources.get(name).copied().unwrap_or_default() +} + +fn optional_string( + params: &Map, + names: &[&str], +) -> Result, AuthError> { + for name in names { + match params.get(*name) { + None | Some(Value::Null) => continue, + Some(Value::String(value)) if value.trim().is_empty() => continue, + Some(Value::String(value)) => return Ok(Some(value.clone())), + Some(_) => { + return Err(AuthError::Configuration( + AuthConfigurationError::InvalidFieldType(names[0].to_string()), + )); + } + } + } + Ok(None) +} + +fn non_empty_env(env_lookup: &dyn Fn(&str) -> Option, name: &str) -> Option { + env_lookup(name) + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) +} + +fn auth_acquisition_error(error: gcp_auth::Error) -> AuthError { + AuthError::VertexTokenAcquisition(error.to_string()) +} + +#[cfg(test)] +mod tests { + use std::sync::atomic::{AtomicUsize, Ordering}; + + use serde_json::json; + + use super::*; + + struct FakeProvider { + calls: Arc, + } + + impl VertexTokenSource for FakeProvider { + fn project_id(&self) -> VertexAuthFuture<'_, String> { + self.calls.fetch_add(1, Ordering::SeqCst); + Box::pin(async { Ok("adc-project".into()) }) + } + + fn token(&self) -> VertexAuthFuture<'_, String> { + self.calls.fetch_add(1, Ordering::SeqCst); + Box::pin(async { Ok("adc-token".into()) }) + } + } + + struct FakeLoader { + loads: Arc, + provider: Arc, + } + + impl VertexProviderLoader for FakeLoader { + fn load( + &self, + _source: CredentialSource, + ) -> VertexAuthFuture<'_, Arc> { + let loads = self.loads.clone(); + let provider = self.provider.clone(); + Box::pin(async move { + loads.fetch_add(1, Ordering::SeqCst); + Ok(provider) + }) + } + } + + fn config(value: Value) -> VertexConfig { + VertexConfig::from_sourced_optional_params(value.as_object().unwrap(), &BTreeMap::new()) + .unwrap() + } + + fn auth(calls: Arc, loads: Arc) -> VertexAuth { + let provider: Arc = Arc::new(FakeProvider { calls }); + VertexAuth::new(Arc::new(FakeLoader { loads, provider })) + } + + #[test] + fn config_is_typed_and_secrets_are_redacted() { + let config = config(json!({ + "vertex_credentials":{"private_key":"secret-key"}, + "vertex_project":"project-1", + "vertex_location":"europe-west4" + })); + assert_eq!(config.project_id(), Some("project-1")); + assert_eq!(config.location(), Some("europe-west4")); + assert!(!format!("{config:?}").contains("secret-key")); + assert!( + VertexConfig::from_sourced_optional_params( + json!({"vertex_credentials":true}).as_object().unwrap(), + &BTreeMap::new() + ) + .is_err() + ); + } + + #[test] + fn empty_primary_values_fall_back_to_python_aliases() { + let config = config(json!({ + "vertex_credentials": null, + "vertex_ai_credentials": "alias-credentials", + "vertex_project": " ", + "vertex_ai_project": "alias-project", + "vertex_location": null, + "vertex_ai_location": "alias-location" + })); + assert_eq!( + config.credentials.as_ref().unwrap().value().expose(), + "alias-credentials" + ); + assert_eq!(config.project_id(), Some("alias-project")); + assert_eq!(config.location(), Some("alias-location")); + } + + #[test] + fn project_and_location_prefer_input_then_environment() { + let configured = + config(json!({"vertex_project":"input-project","vertex_location":"input-location"})); + let env = |name: &str| Some(format!("env-{name}")); + assert_eq!( + get_vertex_ai_project(&configured, &env).as_deref(), + Some("input-project") + ); + assert_eq!( + get_vertex_ai_location(&configured, &env).as_deref(), + Some("input-location") + ); + let empty = VertexConfig::default(); + assert_eq!( + get_vertex_ai_project(&empty, &|_| Some("env-project".into())).as_deref(), + Some("env-project") + ); + assert_eq!( + get_vertex_ai_location(&empty, &|name| (name == VERTEX_LOCATION_ENV) + .then(|| "fallback-location".into())) + .as_deref(), + Some("fallback-location") + ); + } + + #[test] + fn credential_discovery_prefers_input_then_environment_then_adc() { + let params = json!({"vertex_credentials":"input-json"}); + let sources = BTreeMap::from([("vertex_credentials".to_string(), InputSource::Request)]); + let configured = + VertexConfig::from_sourced_optional_params(params.as_object().unwrap(), &sources) + .unwrap(); + assert!( + matches!(credential_source(&configured, &|_| Some("environment-value".into())), CredentialSource::Inline(value) if value.expose() == "input-json") + ); + let empty = VertexConfig::default(); + assert!( + matches!(credential_source(&empty, &|name| (name == VERTEXAI_CREDENTIALS_ENV).then(|| "environment-json".into())), CredentialSource::Trusted(value) if value.expose() == "environment-json") + ); + assert!( + matches!(credential_source(&empty, &|name| (name == GOOGLE_APPLICATION_CREDENTIALS_ENV).then(|| "adc.json".into())), CredentialSource::ApplicationCredentials(path) if path == "adc.json") + ); + assert!(matches!( + credential_source(&empty, &|_| None), + CredentialSource::Adc + )); + assert_ne!( + CredentialSource::Inline(SecretValue::new("same-value")).cache_key(), + CredentialSource::Trusted(SecretValue::new("same-value")).cache_key() + ); + } + + #[test] + fn request_credentials_require_canonical_token_endpoint() { + assert!( + validate_request_credentials(r#"{"token_uri":"https://oauth2.googleapis.com/token"}"#) + .is_ok() + ); + assert!(matches!( + validate_request_credentials(r#"{"token_uri":"http://127.0.0.1/token"}"#), + Err(AuthError::Configuration( + AuthConfigurationError::RequestVertexTokenEndpoint + )) + )); + assert!(matches!( + validate_request_credentials("{}"), + Err(AuthError::Configuration( + AuthConfigurationError::RequestVertexTokenEndpoint + )) + )); + } + + #[tokio::test] + async fn explicit_token_and_header_do_not_acquire_adc() { + let loads = Arc::new(AtomicUsize::new(0)); + let auth = auth(Arc::new(AtomicUsize::new(0)), loads.clone()); + let configured = config(json!({"vertex_project":"project-1"})); + let explicit = auth + .validate_environment(Vec::new(), Some("access-token"), &configured, &|_| None) + .await + .unwrap(); + assert_eq!(explicit.headers[0].1, "Bearer access-token"); + let existing = auth + .validate_environment( + vec![("authorization".into(), "Bearer existing".into())], + None, + &configured, + &|_| None, + ) + .await + .unwrap(); + assert_eq!(existing.headers[0].1, "Bearer existing"); + assert_eq!(loads.load(Ordering::SeqCst), 0); + } + + #[tokio::test] + async fn provider_is_reused_across_authentication_calls() { + let calls = Arc::new(AtomicUsize::new(0)); + let loads = Arc::new(AtomicUsize::new(0)); + let auth = auth(calls.clone(), loads.clone()); + for _ in 0..2 { + let environment = auth + .validate_environment(Vec::new(), None, &VertexConfig::default(), &|_| None) + .await + .unwrap(); + assert_eq!(environment.project_id, "adc-project"); + assert_eq!(environment.headers[0].1, "Bearer adc-token"); + } + assert_eq!(loads.load(Ordering::SeqCst), 1); + assert_eq!(calls.load(Ordering::SeqCst), 4); + } +} diff --git a/litellm-rust/crates/core/src/constants.rs b/litellm-rust/crates/core/src/constants.rs index 108d2a48e30..9469d379462 100644 --- a/litellm-rust/crates/core/src/constants.rs +++ b/litellm-rust/crates/core/src/constants.rs @@ -43,6 +43,23 @@ pub const EMPTY_TEXT_PLACEHOLDER: &str = "[System: Empty message content sanitised to satisfy protocol]"; pub const FUNCTION_TRACE_TARGET: &str = "litellm::function_trace"; + +pub(crate) const MEDIA_CONNECT_TIMEOUT_SECS: u64 = 10; + pub(crate) const OCR_HTTP_TIMEOUT_SECS: u64 = 600; pub(crate) const OCR_CONNECT_TIMEOUT_SECS: u64 = 10; +pub(crate) const OCR_INLINE_MAX_BYTES: usize = 50 * 1024 * 1024; +pub(crate) const OCR_DOWNLOAD_MAX_BYTES: u64 = 50 * 1024 * 1024; +pub(crate) const OCR_MAX_FETCH_REDIRECTS: usize = 10; +pub(crate) const OCR_POLL_TIMEOUT_SECS: u64 = 120; +pub(crate) const OCR_POLL_RETRY_SECS: u64 = 2; +pub(crate) const AZURE_DI_API_VERSION: &str = "2024-11-30"; +pub(crate) const AZURE_DI_SUBSCRIPTION_HEADER: &str = "Ocp-Apim-Subscription-Key"; +pub(crate) const AZURE_DI_DEFAULT_DPI: i64 = 96; +pub(crate) const AZURE_DI_DEFAULT_WIDTH: f64 = 8.5; +pub(crate) const AZURE_DI_DEFAULT_HEIGHT: f64 = 11.0; +pub(crate) const REDUCTO_API_BASE: &str = "https://platform.reducto.ai"; +pub(crate) const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY"; +pub(crate) const REDUCTO_ID_PREFIX: &str = "reducto://"; +pub(crate) const AZURE_AI_OCR_PATH: &str = "/providers/mistral/azure/ocr"; pub(crate) const MISTRAL_OCR_API_BASE: &str = "https://api.mistral.ai/v1"; diff --git a/litellm-rust/crates/core/src/error.rs b/litellm-rust/crates/core/src/error.rs index 0382314057f..fa4a9d36e03 100644 --- a/litellm-rust/crates/core/src/error.rs +++ b/litellm-rust/crates/core/src/error.rs @@ -21,6 +21,18 @@ pub enum Error { "Missing {provider} API Key - A call is being made to {provider} but no key is set either in the environment variables or via params" )] MissingApiKey { provider: &'static str }, + #[error( + "invalid authentication configuration: Missing Azure AI credentials - set AZURE_AI_API_KEY or configure Entra ID" + )] + MissingAzureAiCredentials, + #[error( + "invalid authentication configuration: Missing Azure Document Intelligence credentials - set AZURE_DOCUMENT_INTELLIGENCE_API_KEY or configure Entra ID" + )] + MissingAzureDocumentIntelligenceCredentials, + #[error( + "Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()" + )] + MissingReductoApiKey, #[error("upstream request failed with status {status}: {body}")] Http { status: u16, body: String }, #[error("upstream network error: {0}")] @@ -40,6 +52,28 @@ pub enum Error { Unsupported(&'static str), } +#[derive(Debug, ThisError)] +pub(crate) enum MediaError { + #[error("media URL rejected by network policy")] + BlockedUrl, + #[error("media download is disabled")] + DownloadDisabled, + #[error("media download exceeds the maximum size")] + DownloadTooLarge, + #[error("too many redirects while fetching media")] + TooManyRedirects, + #[error("media redirect is missing a Location header")] + MissingRedirectLocation, + #[error("invalid media redirect")] + InvalidRedirect, + #[error("media download failed with status {0}")] + Http(u16), + #[error("media download timed out")] + Timeout, + #[error("{0}")] + Transport(#[from] TransportError), +} + #[derive(Clone, Debug, ThisError, PartialEq, Eq)] pub enum TransportError { #[error("upstream request failed with status {status}: {body}")] @@ -93,6 +127,15 @@ impl From for Error { } } +impl From for Error { + fn from(error: crate::AuthError) -> Self { + match error { + crate::AuthError::MissingApiKey { provider } => Self::MissingApiKey { provider }, + error => Self::Auth(error.to_string()), + } + } +} + pub fn json_type_name(value: &serde_json::Value) -> &'static str { match value { serde_json::Value::Null => "null", @@ -108,6 +151,14 @@ pub fn json_type_name(value: &serde_json::Value) -> &'static str { mod transport_tests { use super::*; + #[test] + fn missing_auth_key_preserves_provider_in_public_error() { + assert_eq!( + Error::from(crate::AuthError::MissingApiKey { provider: "Vertex" }), + Error::MissingApiKey { provider: "Vertex" } + ); + } + #[tokio::test] async fn transport_errors_remove_urls_and_keep_dispatch_context() { let error = reqwest::Client::builder() diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 3a0896a4d5a..0b3573deab2 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,10 +1,12 @@ pub mod audio_transcription; +pub mod auth; pub mod caching; pub mod call_lifecycle; pub mod chat_completions; pub mod constants; pub mod error; pub mod http_utils; +mod media; pub mod messages; #[cfg(any(feature = "observability", test))] pub mod observability; @@ -16,4 +18,5 @@ pub mod router; pub mod routing_utils; mod url_utils; +pub use auth::AuthError; pub use error::Error; diff --git a/litellm-rust/crates/core/src/media.rs b/litellm-rust/crates/core/src/media.rs new file mode 100644 index 00000000000..5f9a43794c2 --- /dev/null +++ b/litellm-rust/crates/core/src/media.rs @@ -0,0 +1,528 @@ +use std::future::Future; +use std::io; +use std::net::{IpAddr, SocketAddr}; +use std::pin::Pin; +use std::sync::Arc; +use std::time::Duration; + +use reqwest::Url; +use reqwest::dns::{Addrs, Name, Resolve, Resolving}; + +use crate::constants::MEDIA_CONNECT_TIMEOUT_SECS; +use crate::error::{MediaError, TransportError}; + +#[derive(Clone)] +pub(crate) struct MediaFetcher { + client: reqwest::Client, + address_resolver: Arc, + allow_private_network: bool, +} + +type AddressResolution<'a> = Pin>> + Send + 'a>>; + +trait AddressResolver: Send + Sync { + fn resolve<'a>(&'a self, host: &'a str, port: u16) -> AddressResolution<'a>; +} + +#[derive(Clone, Copy)] +pub(crate) struct DownloadPolicy { + pub(crate) timeout: Duration, + pub(crate) max_bytes: u64, + pub(crate) max_redirects: usize, +} + +#[derive(Debug)] +pub(crate) struct DownloadedMedia { + pub(crate) bytes: Vec, + pub(crate) content_type: String, +} + +impl MediaFetcher { + pub(crate) fn new() -> Result { + Self::with_resolvers(Arc::new(PublicDnsResolver), Arc::new(SystemAddressResolver)) + } + + fn with_resolvers( + transport_resolver: Arc, + address_resolver: Arc, + ) -> Result + where + R: Resolve + 'static, + { + let client = reqwest::Client::builder() + .connect_timeout(Duration::from_secs(MEDIA_CONNECT_TIMEOUT_SECS)) + .redirect(reqwest::redirect::Policy::none()) + .no_proxy() + .dns_resolver(transport_resolver) + .build()?; + Ok(Self { + client, + address_resolver, + allow_private_network: false, + }) + } + + #[cfg(test)] + pub(crate) fn for_test(client: reqwest::Client) -> Self { + Self { + client, + address_resolver: Arc::new(AllowPrivateResolver), + allow_private_network: true, + } + } + + pub(crate) async fn fetch( + &self, + url: Url, + policy: DownloadPolicy, + ) -> Result { + if policy.max_bytes == 0 { + return Err(MediaError::DownloadDisabled); + } + tokio::time::timeout(policy.timeout, self.fetch_before_deadline(url, policy)) + .await + .map_err(|_| MediaError::Timeout)? + } + + async fn fetch_before_deadline( + &self, + mut url: Url, + policy: DownloadPolicy, + ) -> Result { + let mut redirects_followed = 0; + loop { + self.validate_url(&url).await?; + let mut response = self + .client + .get(url.clone()) + .send() + .await + .map_err(TransportError::from)?; + if response.status().is_redirection() { + if redirects_followed == policy.max_redirects { + return Err(MediaError::TooManyRedirects); + } + let location = response + .headers() + .get(reqwest::header::LOCATION) + .and_then(|value| value.to_str().ok()) + .ok_or(MediaError::MissingRedirectLocation)?; + url = url + .join(location) + .map_err(|_| MediaError::InvalidRedirect)?; + redirects_followed += 1; + continue; + } + if !response.status().is_success() { + return Err(MediaError::Http(response.status().as_u16())); + } + enforce_download_size(response.content_length().unwrap_or(0), policy.max_bytes)?; + let content_type = response + .headers() + .get(reqwest::header::CONTENT_TYPE) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.split(';').next()) + .map(str::trim) + .filter(|value| !value.is_empty()) + .unwrap_or("application/octet-stream") + .to_string(); + let mut bytes = Vec::new(); + while let Some(chunk) = response.chunk().await.map_err(TransportError::from)? { + enforce_download_size(bytes.len() as u64 + chunk.len() as u64, policy.max_bytes)?; + bytes.extend_from_slice(&chunk); + } + return Ok(DownloadedMedia { + bytes, + content_type, + }); + } + } + + async fn validate_url(&self, url: &Url) -> Result<(), MediaError> { + if !matches!(url.scheme(), "http" | "https") + || !url.username().is_empty() + || url.password().is_some() + { + return Err(MediaError::BlockedUrl); + } + let host = url.host_str().ok_or(MediaError::BlockedUrl)?; + if self.allow_private_network { + return Ok(()); + } + if let Ok(ip) = host.parse::() { + return (!is_blocked_ip(ip)) + .then_some(()) + .ok_or(MediaError::BlockedUrl); + } + let port = url.port_or_known_default().ok_or(MediaError::BlockedUrl)?; + let addresses = self + .address_resolver + .resolve(host, port) + .await + .map_err(|error| TransportError::Network(error.to_string()))?; + validate_addresses(&addresses) + } +} + +fn enforce_download_size(length: u64, max_bytes: u64) -> Result<(), MediaError> { + if length > max_bytes { + return Err(MediaError::DownloadTooLarge); + } + Ok(()) +} + +fn validate_addresses(addresses: &[SocketAddr]) -> Result<(), MediaError> { + if addresses.is_empty() || addresses.iter().any(|address| is_blocked_ip(address.ip())) { + return Err(MediaError::BlockedUrl); + } + Ok(()) +} + +fn is_blocked_ip(ip: IpAddr) -> bool { + match ip { + IpAddr::V4(ip) => { + let [first, second, third, _] = ip.octets(); + first == 0 + || first == 10 + || first == 127 + || (first == 100 && (64..=127).contains(&second)) + || (first == 169 && second == 254) + || (first == 172 && (16..=31).contains(&second)) + || (first == 192 && second == 0 && (third == 0 || third == 2)) + || (first == 192 && second == 168) + || (first == 192 && second == 88 && third == 99) + || (first == 198 && (second == 18 || second == 19)) + || (first == 198 && second == 51 && third == 100) + || (first == 203 && second == 0 && third == 113) + || first >= 224 + } + IpAddr::V6(ip) => { + let segments = ip.segments(); + ip.is_loopback() + || ip.is_unspecified() + || ip.is_multicast() + || (segments[0] & 0xfe00) == 0xfc00 + || (segments[0] & 0xffc0) == 0xfe80 + || (segments[0] & 0xffc0) == 0xfec0 + || (segments[0] == 0x2001 && segments[1] == 0x0db8) + || ip + .to_ipv4_mapped() + .or_else(|| ip.to_ipv4()) + .map(|ipv4| is_blocked_ip(IpAddr::V4(ipv4))) + .unwrap_or(false) + } + } +} + +#[derive(Default)] +struct PublicDnsResolver; + +struct SystemAddressResolver; + +impl AddressResolver for SystemAddressResolver { + fn resolve<'a>(&'a self, host: &'a str, port: u16) -> AddressResolution<'a> { + Box::pin(async move { + Ok(tokio::net::lookup_host((host, port)) + .await? + .collect::>()) + }) + } +} + +#[cfg(test)] +struct AllowPrivateResolver; + +#[cfg(test)] +impl AddressResolver for AllowPrivateResolver { + fn resolve<'a>(&'a self, _host: &'a str, port: u16) -> AddressResolution<'a> { + Box::pin(async move { Ok(vec![SocketAddr::from(([8, 8, 8, 8], port))]) }) + } +} + +impl Resolve for PublicDnsResolver { + fn resolve(&self, name: Name) -> Resolving { + let host = name.as_str().to_string(); + Box::pin(async move { + let addresses = tokio::net::lookup_host((host.as_str(), 0)) + .await + .map_err(|error| Box::new(error) as Box)? + .collect::>(); + validate_addresses(&addresses).map_err(|_| { + Box::new(io::Error::other("destination rejected by network policy")) + as Box + })?; + Ok(Box::new(addresses.into_iter()) as Addrs) + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashSet; + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + async fn serve(response: &'static [u8]) -> (Url, tokio::task::JoinHandle<()>) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let address = listener.local_addr().expect("listener has address"); + let task = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let mut request = [0_u8; 1024]; + let bytes_read = socket.read(&mut request).await.expect("reads request"); + assert!(bytes_read > 0); + socket.write_all(response).await.expect("writes response"); + }); + ( + Url::parse(&format!("http://{address}/document")).expect("valid test URL"), + task, + ) + } + + async fn serve_named( + host: &str, + responses: Vec<&'static [u8]>, + ) -> (Url, tokio::task::JoinHandle>, SocketAddr) { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let address = listener.local_addr().expect("listener has address"); + let task = tokio::spawn(async move { + let mut requests = Vec::with_capacity(responses.len()); + for response in responses { + let (mut socket, _) = listener.accept().await.expect("accepts request"); + let mut request = [0_u8; 4096]; + let bytes_read = socket.read(&mut request).await.expect("reads request"); + requests.push(String::from_utf8_lossy(&request[..bytes_read]).into_owned()); + socket.write_all(response).await.expect("writes response"); + } + requests + }); + ( + Url::parse(&format!("http://{host}:{}/document", address.port())) + .expect("valid test URL"), + task, + address, + ) + } + + struct LoopbackDnsResolver(SocketAddr); + + impl Resolve for LoopbackDnsResolver { + fn resolve(&self, _name: Name) -> Resolving { + let address = self.0; + Box::pin(async move { Ok(Box::new(vec![address].into_iter()) as Addrs) }) + } + } + + struct TestAddressResolver { + blocked_hosts: HashSet<&'static str>, + } + + impl AddressResolver for TestAddressResolver { + fn resolve<'a>(&'a self, host: &'a str, port: u16) -> AddressResolution<'a> { + let blocked = self.blocked_hosts.contains(host); + Box::pin(async move { + let ip = if blocked { + IpAddr::from([127, 0, 0, 1]) + } else { + IpAddr::from([8, 8, 8, 8]) + }; + Ok(vec![SocketAddr::new(ip, port)]) + }) + } + } + + fn policy_checked_fetcher( + address: SocketAddr, + blocked_hosts: HashSet<&'static str>, + ) -> MediaFetcher { + MediaFetcher::with_resolvers( + Arc::new(LoopbackDnsResolver(address)), + Arc::new(TestAddressResolver { blocked_hosts }), + ) + .expect("test fetcher builds") + } + + fn policy(max_bytes: u64, max_redirects: usize) -> DownloadPolicy { + DownloadPolicy { + timeout: Duration::from_secs(1), + max_bytes, + max_redirects, + } + } + + #[test] + fn blocks_non_public_addresses() { + for address in [ + "0.0.0.1", + "10.0.0.1", + "100.64.0.1", + "127.0.0.1", + "169.254.1.1", + "172.16.0.1", + "192.168.0.1", + "198.18.0.1", + "198.51.100.1", + "203.0.113.1", + "224.0.0.1", + "::1", + "fc00::1", + "fe80::1", + "2001:db8::1", + "::ffff:127.0.0.1", + ] { + assert!(is_blocked_ip(address.parse().expect("valid test address"))); + } + assert!(!is_blocked_ip( + "8.8.8.8".parse().expect("valid public address") + )); + } + + #[tokio::test] + async fn fetches_exact_limit_and_normalizes_content_type() { + let (url, server) = serve( + b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf; charset=binary\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc", + ) + .await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("test client builds"); + let media = MediaFetcher::for_test(client) + .fetch(url, policy(3, 0)) + .await + .expect("download succeeds at exact limit"); + server.await.expect("server completes"); + assert_eq!(media.bytes, b"abc"); + assert_eq!(media.content_type, "application/pdf"); + } + + #[tokio::test] + async fn rejects_declared_oversize_body() { + let (url, server) = serve( + b"HTTP/1.1 200 OK\r\nContent-Type: application/pdf\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc", + ) + .await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("test client builds"); + let error = MediaFetcher::for_test(client) + .fetch(url, policy(2, 0)) + .await + .expect_err("oversize body is rejected"); + server.await.expect("server completes"); + assert!(matches!(error, MediaError::DownloadTooLarge)); + } + + #[tokio::test] + async fn rejects_streamed_oversize_body() { + let (url, server) = serve( + b"HTTP/1.1 200 OK\r\nTransfer-Encoding: chunked\r\nConnection: close\r\n\r\n2\r\nab\r\n2\r\ncd\r\n0\r\n\r\n", + ) + .await; + let client = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("test client builds"); + let error = MediaFetcher::for_test(client) + .fetch(url, policy(3, 0)) + .await + .expect_err("stream crossing limit is rejected"); + server.await.expect("server completes"); + assert!(matches!(error, MediaError::DownloadTooLarge)); + } + + #[tokio::test] + async fn follows_allowed_redirects_and_revalidates_each_destination() { + let (url, server, address) = serve_named( + "public.test", + vec![ + b"HTTP/1.1 302 Found\r\nLocation: /final\r\nContent-Length: 0\r\n\r\n", + b"HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok", + ], + ) + .await; + let media = policy_checked_fetcher(address, HashSet::new()) + .fetch(url, policy(2, 1)) + .await + .expect("redirected fetch succeeds"); + let requests = server.await.expect("server completes"); + assert_eq!(requests.len(), 2); + assert!(requests[1].starts_with("GET /final ")); + assert_eq!(media.bytes, b"ok"); + } + + #[tokio::test] + async fn blocks_redirected_private_destination_before_second_request() { + let (url, server, address) = serve_named( + "public.test", + vec![b"HTTP/1.1 302 Found\r\nLocation: http://blocked.test/document\r\nContent-Length: 0\r\n\r\n"], + ) + .await; + let error = policy_checked_fetcher(address, HashSet::from(["blocked.test"])) + .fetch(url, policy(10, 1)) + .await + .expect_err("private redirect is rejected"); + let requests = server.await.expect("server completes"); + assert_eq!(requests.len(), 1); + assert!(matches!(error, MediaError::BlockedUrl)); + } + + #[tokio::test] + async fn enforces_total_timeout() { + let listener = TcpListener::bind("127.0.0.1:0") + .await + .expect("test listener binds"); + let address = listener.local_addr().expect("listener has address"); + let server = tokio::spawn(async move { + let (_socket, _) = listener.accept().await.expect("accepts request"); + tokio::time::sleep(Duration::from_millis(100)).await; + }); + let url = Url::parse(&format!("http://public.test:{}/document", address.port())) + .expect("valid test URL"); + let error = policy_checked_fetcher(address, HashSet::new()) + .fetch( + url, + DownloadPolicy { + timeout: Duration::from_millis(20), + max_bytes: 10, + max_redirects: 0, + }, + ) + .await + .expect_err("fetch times out"); + server.await.expect("server completes"); + assert!(matches!(error, MediaError::Timeout)); + } + + #[tokio::test] + async fn document_client_does_not_send_ambient_credentials() { + let (url, server, address) = serve_named( + "public.test", + vec![b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nConnection: close\r\n\r\nok"], + ) + .await; + policy_checked_fetcher(address, HashSet::new()) + .fetch(url, policy(2, 0)) + .await + .expect("fetch succeeds"); + let requests = server.await.expect("server completes"); + assert!(!requests[0].to_ascii_lowercase().contains("authorization:")); + assert!(!requests[0].to_ascii_lowercase().contains("api-key:")); + } + + #[tokio::test] + async fn rejects_url_credentials_before_network_access() { + let fetcher = MediaFetcher::new().expect("media fetcher builds"); + let url = + Url::parse("https://user:password@8.8.8.8/document").expect("credentialed URL parses"); + assert!(matches!( + fetcher.validate_url(&url).await, + Err(MediaError::BlockedUrl) + )); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs new file mode 100644 index 00000000000..71ca69ddc58 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/mod.rs @@ -0,0 +1,214 @@ +use super::super::OcrAdapter; +use crate::Error; +use crate::auth::{InputSource, Sourced}; +use crate::constants::{AZURE_DI_API_VERSION, AZURE_DI_SUBSCRIPTION_HEADER}; +use crate::ocr::OcrClient; +use crate::ocr::codecs::document_intelligence::{ + self, AzureDocumentIntelligenceOperation, DocumentIntelligenceParams, +}; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{credential_env, transform_request_body}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrResponseFormat}; +use crate::ocr::wire::DecodedOcrResponse; +use crate::providers::azure_ai::auth::AzureAuthInputs; +use crate::url_utils::ApiUrl; + +mod polling; + +const AZURE_DI_API_KEY_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"; +const AZURE_DI_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT"; + +#[derive(Clone, Debug)] +pub(crate) struct AzureDocumentIntelligenceAdapter; + +impl OcrAdapter for AzureDocumentIntelligenceAdapter { + type ProviderResponse = AzureDocumentIntelligenceOperation; + const PROVIDER: OcrProvider = OcrProvider::AzureAi; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let params = map_ocr_params(request)?; + let config = AzureAuthInputs::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)?; + let headers = validate_environment(&request.connection, &config, &credential_env).await?; + let endpoint = nonblank(request.connection.api_base.clone()) + .or_else(|| nonblank(credential_env(AZURE_DI_ENDPOINT_ENV))) + .ok_or_else(|| Error::Auth("Missing Azure Document Intelligence API Base - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT or pass api_base".into()))?; + let url = get_complete_url(&endpoint, &request.model, ¶ms)?; + let body = document_intelligence::transform_ocr_request(request.document.clone())?; + transform_request_body(client, request, &url, &headers, body, |_| Ok(())).await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + document_intelligence::transform_ocr_response(&request.model, response) + } + + async fn read_response( + &self, + client: &OcrClient, + response: reqwest::Response, + url: &str, + headers: &[(String, String)], + request: &LiteLLMOcrRequest, + ) -> Result, OcrError> { + polling::read_operation_response( + client.polling_http(), + response, + url, + headers, + &request.connection, + request.response_format()? == OcrResponseFormat::Native, + ) + .await + } +} + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +fn map_ocr_params( + request: &LiteLLMOcrRequest, +) -> Result { + let params = document_intelligence::decode_input_params( + request.optional_params.clone(), + "optional_params", + )?; + let crate::ocr::prepare::ParsedProviderParams { + known: params, + extra_params: _extra_params, + } = params; + document_intelligence::map_ocr_params(params) +} + +fn get_complete_url( + endpoint: &str, + model: &str, + params: &DocumentIntelligenceParams, +) -> Result { + let model = format!("{}:analyze", model_id(model)?); + ApiUrl::parse(endpoint) + .and_then(|url| url.complete_path(&["documentintelligence", "documentModels", &model])) + .map(|url| { + url.append_query_pairs( + [("api-version", AZURE_DI_API_VERSION)] + .into_iter() + .chain(params.pages.iter().map(|pages| ("pages", pages.as_str()))) + .chain( + params + .features + .iter() + .map(|features| ("features", features.as_str())), + ), + ) + .into_string() + }) + .map_err(|_| OcrRequestError::RequestField { + path: "api_base".into(), + }) + .map_err(OcrError::from) +} + +async fn validate_environment( + connection: &OcrConnection, + config: &AzureAuthInputs, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result, OcrError> { + if crate::http_utils::has_header(&connection.extra_headers, "authorization") + || crate::http_utils::has_header(&connection.extra_headers, AZURE_DI_SUBSCRIPTION_HEADER) + { + super::validate_destination(connection, connection.extra_headers_source)?; + return Ok(connection.extra_headers.clone()); + } + let key = nonblank(connection.api_key.clone()) + .map(|value| Sourced::new(value, connection.api_key_source)) + .or_else(|| { + nonblank(env_lookup(AZURE_DI_API_KEY_ENV)) + .map(|value| Sourced::new(value, InputSource::Environment)) + }); + if let Some(key) = key { + super::validate_destination(connection, key.source())?; + return Ok( + std::iter::once((AZURE_DI_SUBSCRIPTION_HEADER.into(), key.into_value())) + .chain(connection.extra_headers.clone()) + .collect(), + ); + } + let token = super::resolve_entra(config, env_lookup) + .await? + .ok_or(Error::MissingAzureDocumentIntelligenceCredentials)?; + super::validate_destination(connection, token.source())?; + Ok( + std::iter::once(("Authorization".into(), format!("Bearer {}", token.value()))) + .chain(connection.extra_headers.clone()) + .collect(), + ) +} + +fn model_id(model: &str) -> Result<&str, OcrRequestError> { + let model = model.rsplit('/').next().unwrap_or(model); + if matches!(model, "." | "..") { + return Err(OcrRequestError::DotModel); + } + Ok(model) +} + +fn nonblank(value: Option) -> Option { + value + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[tokio::test] + async fn request_endpoint_cannot_receive_environment_key() { + let connection = OcrConnection { + api_base: Some("https://request.example".into()), + api_base_source: InputSource::Request, + ..Default::default() + }; + + let error = validate_environment(&connection, &Default::default(), &|name| { + (name == AZURE_DI_API_KEY_ENV).then(|| "environment-key".into()) + }) + .await + .unwrap_err(); + + assert!( + error + .to_string() + .contains("request-controlled Azure endpoint") + ); + } + + #[tokio::test] + async fn request_endpoint_accepts_request_owned_key() { + let connection = OcrConnection { + api_key: Some("request-key".into()), + api_key_source: InputSource::Request, + api_base: Some("https://request.example".into()), + api_base_source: InputSource::Request, + ..Default::default() + }; + + let headers = validate_environment(&connection, &Default::default(), &|_| None) + .await + .unwrap(); + + assert_eq!( + headers[0], + (AZURE_DI_SUBSCRIPTION_HEADER.into(), "request-key".into()) + ); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs new file mode 100644 index 00000000000..1bddea0da4f --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/document_intelligence/polling.rs @@ -0,0 +1,100 @@ +use std::time::Duration; + +use reqwest::Url; +use tokio::time::Instant; + +use crate::constants::{AZURE_DI_SUBSCRIPTION_HEADER, OCR_POLL_RETRY_SECS}; +use crate::ocr::client::read_json_response; +use crate::ocr::codecs::document_intelligence::{ + AzureDocumentIntelligenceOperation, OperationStatus, +}; +use crate::ocr::error::{OcrError, OcrPollingError, OcrResponseError}; +use crate::ocr::types::OcrConnection; +use crate::ocr::wire::DecodedOcrResponse; + +pub(super) async fn read_operation_response( + http_client: &reqwest::Client, + response: reqwest::Response, + original_url: &str, + headers: &[(String, String)], + connection: &OcrConnection, + native: bool, +) -> Result, OcrError> { + if response.status() != reqwest::StatusCode::ACCEPTED { + return read_json_response(response, native).await; + } + let location = response + .headers() + .get("operation-location") + .and_then(|value| value.to_str().ok()) + .ok_or(OcrPollingError::PollLocation)?; + let original = Url::parse(original_url).map_err(|_| OcrPollingError::PollOrigin)?; + let operation = Url::parse(location).map_err(|_| OcrPollingError::PollOrigin)?; + if original.origin() != operation.origin() + || !operation.username().is_empty() + || operation.password().is_some() + { + return Err(OcrPollingError::PollOrigin.into()); + } + poll_operation(http_client, operation, headers, connection, native).await +} + +async fn poll_operation( + http_client: &reqwest::Client, + url: Url, + headers: &[(String, String)], + connection: &OcrConnection, + native: bool, +) -> Result, OcrError> { + let deadline = Instant::now() + .checked_add(connection.poll_timeout) + .ok_or(OcrPollingError::PollTimeout)?; + loop { + let remaining = deadline + .checked_duration_since(Instant::now()) + .filter(|remaining| !remaining.is_zero()) + .ok_or(OcrPollingError::PollTimeout)?; + let builder = http_client + .get(url.clone()) + .timeout(remaining.min(connection.timeout)); + let builder = crate::http_utils::with_headers( + builder, + headers, + crate::http_utils::HeaderPolicy::Only(&[AZURE_DI_SUBSCRIPTION_HEADER, "authorization"]), + ); + let response = tokio::time::timeout_at(deadline, crate::http_utils::http_request(builder)) + .await + .map_err(|_| OcrPollingError::PollTimeout)? + .map_err(crate::error::TransportError::from)?; + let retry = response + .headers() + .get(reqwest::header::RETRY_AFTER) + .and_then(|value| value.to_str().ok()) + .and_then(|value| value.parse::().ok()) + .unwrap_or(OCR_POLL_RETRY_SECS) + .max(1); + let decoded = tokio::time::timeout_at( + deadline, + read_json_response::(response, native), + ) + .await + .map_err(|_| OcrPollingError::PollTimeout)??; + match &decoded.data.status { + Some(OperationStatus::Succeeded) => return Ok(decoded), + Some(OperationStatus::Running | OperationStatus::NotStarted) => { + tokio::time::timeout_at(deadline, tokio::time::sleep(Duration::from_secs(retry))) + .await + .map_err(|_| OcrPollingError::PollTimeout)?; + } + status => { + return Err(OcrResponseError::OperationStatus( + status + .as_ref() + .map(ToString::to_string) + .unwrap_or_else(|| "None".into()), + ) + .into()); + } + } + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs new file mode 100644 index 00000000000..3107494d39e --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/mistral.rs @@ -0,0 +1,217 @@ +use super::super::OcrAdapter; +use crate::Error; +use crate::auth::{InputSource, Sourced}; +use crate::constants::AZURE_AI_OCR_PATH; +use crate::ocr::OcrClient; +use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse}; +use crate::ocr::document::{inline_remote_document, validate_inline_document}; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, +}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection}; +use crate::providers::azure_ai::auth::AzureAuthInputs; +use crate::url_utils::ApiUrl; + +const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY"; +const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE"; + +#[derive(Clone, Debug)] +pub(crate) struct AzureMistralAdapter; + +impl OcrAdapter for AzureMistralAdapter { + type ProviderResponse = MistralOcrResponse; + const PROVIDER: OcrProvider = OcrProvider::AzureAi; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params: _extra_params, + } = _prepare_ocr_request::(request)?; + let config = AzureAuthInputs::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)?; + let headers = validate_environment(&request.connection, &config, &credential_env).await?; + let url = get_complete_url(request.connection.api_base.as_deref(), &credential_env)?; + let document = inline_remote_document( + client.document_fetcher(), + request.document.clone(), + &request.connection, + ) + .await?; + let body = mistral::transform_ocr_request(&request.model, document, ¶ms)?; + transform_request_body(client, request, &url, &headers, body, |body| { + validate_inline_document(&body.document) + }) + .await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + mistral::transform_ocr_response(&request.model, response) + } +} + +fn get_complete_url( + api_base: Option<&str>, + env_lookup: &dyn Fn(&str) -> Option, +) -> Result { + let base = nonblank(api_base.map(str::to_string)) + .or_else(|| nonblank(env_lookup(AZURE_AI_API_BASE_ENV))) + .ok_or_else(|| Error::Auth( + "Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter".into(), + ))?; + let path: Vec<&str> = AZURE_AI_OCR_PATH.trim_matches('/').split('/').collect(); + ApiUrl::parse(&base) + .and_then(|url| url.complete_path(&path)) + .map(|url| url.into_string()) + .map_err(|_| { + OcrRequestError::RequestField { + path: "api_base".into(), + } + .into() + }) +} + +async fn validate_environment( + connection: &OcrConnection, + config: &AzureAuthInputs, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result, OcrError> { + if crate::http_utils::has_header(&connection.extra_headers, "authorization") { + super::validate_destination(connection, connection.extra_headers_source)?; + return Ok(connection.extra_headers.clone()); + } + let key = nonblank(connection.api_key.clone()) + .map(|value| Sourced::new(value, connection.api_key_source)) + .or_else(|| { + nonblank(env_lookup(AZURE_AI_API_KEY_ENV)) + .map(|value| Sourced::new(value, InputSource::Environment)) + }); + if let Some(key) = key { + super::validate_destination(connection, key.source())?; + return Ok(bearer_headers(connection, key.value())); + } + let key = super::resolve_entra(config, env_lookup) + .await? + .ok_or(Error::MissingAzureAiCredentials)?; + super::validate_destination(connection, key.source())?; + Ok(bearer_headers(connection, key.value())) +} + +fn bearer_headers(connection: &OcrConnection, key: &str) -> Vec<(String, String)> { + std::iter::once(("Authorization".into(), format!("Bearer {key}"))) + .chain(connection.extra_headers.clone()) + .collect() +} + +fn nonblank(value: Option) -> Option { + value + .map(|value| value.trim().to_string()) + .filter(|value| !value.is_empty()) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn completes_azure_path_and_preserves_query() { + assert_eq!( + get_complete_url(Some("https://example.com/?tenant=a"), &|_| None).unwrap(), + "https://example.com/providers/mistral/azure/ocr?tenant=a" + ); + assert_eq!( + get_complete_url( + Some("https://example.com/providers/mistral/azure/ocr"), + &|_| None + ) + .unwrap(), + "https://example.com/providers/mistral/azure/ocr" + ); + } + + #[tokio::test] + async fn supplied_authorization_precedes_keys() { + let connection = OcrConnection { + api_key: Some("request-key".into()), + extra_headers: vec![("authorization".into(), "Bearer prepared".into())], + ..Default::default() + }; + assert_eq!( + validate_environment(&connection, &Default::default(), &|_| { + Some("environment-key".into()) + }) + .await + .unwrap(), + connection.extra_headers + ); + } + + #[tokio::test] + async fn request_key_precedes_environment_key() { + let connection = OcrConnection { + api_key: Some("request-key".into()), + ..Default::default() + }; + assert_eq!( + validate_environment(&connection, &Default::default(), &|_| { + Some("environment-key".into()) + }) + .await + .unwrap()[0], + ("Authorization".into(), "Bearer request-key".into()) + ); + } + + #[tokio::test] + async fn request_endpoint_cannot_receive_environment_key() { + let connection = OcrConnection { + api_base: Some("https://request.example".into()), + api_base_source: InputSource::Request, + ..Default::default() + }; + + let error = validate_environment(&connection, &Default::default(), &|name| { + (name == AZURE_AI_API_KEY_ENV).then(|| "environment-key".into()) + }) + .await + .unwrap_err(); + + assert!( + error + .to_string() + .contains("request-controlled Azure endpoint") + ); + } + + #[tokio::test] + async fn request_endpoint_accepts_request_owned_key() { + let connection = OcrConnection { + api_key: Some("request-key".into()), + api_key_source: InputSource::Request, + api_base: Some("https://request.example".into()), + api_base_source: InputSource::Request, + ..Default::default() + }; + + let headers = validate_environment(&connection, &Default::default(), &|_| None) + .await + .unwrap(); + + assert_eq!( + headers[0], + ("Authorization".into(), "Bearer request-key".into()) + ); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs new file mode 100644 index 00000000000..9c02a7471c9 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/azure/mod.rs @@ -0,0 +1,49 @@ +mod document_intelligence; +mod mistral; + +use std::sync::OnceLock; + +use crate::Error; +use crate::auth::error::AuthConfigurationError; +use crate::auth::{InputSource, Sourced}; +use crate::ocr::error::OcrError; +use crate::ocr::types::OcrConnection; +use crate::providers::azure_ai::auth::{AzureAuthInputs, AzureAuthService}; + +pub(crate) use document_intelligence::AzureDocumentIntelligenceAdapter; +pub(crate) use mistral::AzureMistralAdapter; + +async fn resolve_entra( + config: &AzureAuthInputs, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result>, Error> { + static SERVICE: OnceLock = OnceLock::new(); + SERVICE + .get_or_init(AzureAuthService::default) + .get_azure_ad_token(config, env_lookup) + .await + .map(|credential| { + credential.map(|credential| { + let source = credential.source(); + let value = credential.value().secret().expose().to_string(); + Sourced::new(value, source) + }) + }) + .map_err(Error::from) +} + +fn validate_destination( + connection: &OcrConnection, + credential_source: InputSource, +) -> Result<(), OcrError> { + if connection.api_base.is_some() + && connection.api_base_source == InputSource::Request + && credential_source != InputSource::Request + { + return Err(Error::from(crate::AuthError::Configuration( + AuthConfigurationError::RequestAzureCredentialDestination, + )) + .into()); + } + Ok(()) +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/mod.rs index 7dfb08b4dc2..9171d11836c 100644 --- a/litellm-rust/crates/core/src/ocr/adapters/mod.rs +++ b/litellm-rust/crates/core/src/ocr/adapters/mod.rs @@ -8,9 +8,15 @@ use super::registry::OcrProvider; use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrResponseFormat}; use super::wire::DecodedOcrResponse; +mod azure; mod mistral; +mod reducto; +mod vertex; +pub(crate) use azure::{AzureDocumentIntelligenceAdapter, AzureMistralAdapter}; pub(crate) use mistral::MistralAdapter; +pub(crate) use reducto::{ReductoLegacyAdapter, ReductoV3Adapter}; +pub(crate) use vertex::{VertexDeepSeekAdapter, VertexMistralAdapter}; /// Converts a complete LiteLLM OCR call to provider HTTP and normalizes its response. pub(crate) trait OcrAdapter: Send + Sync + Sized + 'static { @@ -62,6 +68,12 @@ macro_rules! for_each_ocr_adapter { ($callback:ident) => { $callback! { Mistral, $crate::ocr::adapters::MistralAdapter, $crate::ocr::adapters::MistralAdapter, Mistral; + AzureMistral, $crate::ocr::adapters::AzureMistralAdapter, $crate::ocr::adapters::AzureMistralAdapter, AzureAi; + AzureDocumentIntelligence, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, $crate::ocr::adapters::AzureDocumentIntelligenceAdapter, AzureAi; + ReductoLegacy, $crate::ocr::adapters::ReductoLegacyAdapter, $crate::ocr::adapters::ReductoLegacyAdapter, Reducto; + ReductoV3, $crate::ocr::adapters::ReductoV3Adapter, $crate::ocr::adapters::ReductoV3Adapter, Reducto; + VertexMistral, $crate::ocr::adapters::VertexMistralAdapter, $crate::ocr::adapters::VertexMistralAdapter, VertexAi; + VertexDeepSeek, $crate::ocr::adapters::VertexDeepSeekAdapter, $crate::ocr::adapters::VertexDeepSeekAdapter, VertexAi; } }; } diff --git a/litellm-rust/crates/core/src/ocr/adapters/reducto/legacy.rs b/litellm-rust/crates/core/src/ocr/adapters/reducto/legacy.rs new file mode 100644 index 00000000000..062a0071a34 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/reducto/legacy.rs @@ -0,0 +1,45 @@ +use super::super::OcrAdapter; +use crate::ocr::OcrClient; +use crate::ocr::codecs::reducto::{self, ReductoLegacyParams, ReductoResponse}; +use crate::ocr::error::{OcrError, OcrResponseError}; +use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env, + guardrail_document, merge_extra_params, +}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; + +#[derive(Clone, Debug)] +pub(crate) struct ReductoLegacyAdapter; + +impl OcrAdapter for ReductoLegacyAdapter { + type ProviderResponse = ReductoResponse; + const PROVIDER: OcrProvider = OcrProvider::Reducto; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params, + } = _prepare_ocr_request::(request)?; + let headers = super::validate_environment(&request.connection, &credential_env)?; + let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?; + let document = guardrail_document(request, &url).await?; + let document = + super::prepare_document(client, document, &request.connection, &headers).await?; + let body = reducto::transform_legacy_ocr_request(&request.model, document, ¶ms)?; + let body = merge_extra_params(&body, extra_params)?; + build_http_request(client, request, &url, &headers, &body) + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + reducto::transform_ocr_response(&request.model, response) + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/reducto/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/reducto/mod.rs new file mode 100644 index 00000000000..7621d0d326a --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/reducto/mod.rs @@ -0,0 +1,148 @@ +mod legacy; +mod v3; + +use crate::Error; +use crate::constants::{REDUCTO_API_BASE, REDUCTO_API_KEY_ENV, REDUCTO_ID_PREFIX}; +use crate::ocr::document::InlineDocument; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::types::{OcrConnection, OcrDocument}; +use crate::url_utils::ApiUrl; + +pub(crate) use legacy::ReductoLegacyAdapter; +pub(crate) use v3::ReductoV3Adapter; + +pub(super) fn get_complete_url(api_base: Option<&str>, path: &str) -> Result { + let base = api_base + .map(str::trim) + .filter(|base| !base.is_empty()) + .unwrap_or(REDUCTO_API_BASE); + ApiUrl::parse(base) + .and_then(|url| url.complete_path(&[path])) + .map(|url| url.into_string()) + .map_err(|_| { + OcrRequestError::RequestField { + path: "api_base".into(), + } + .into() + }) +} + +pub(super) fn validate_environment( + connection: &OcrConnection, + env_lookup: &(dyn Fn(&str) -> Option + Sync), +) -> Result, OcrError> { + if crate::http_utils::has_header(&connection.extra_headers, "authorization") { + return Ok(connection.extra_headers.clone()); + } + let api_key = connection + .api_key + .as_deref() + .map(str::trim) + .filter(|key| !key.is_empty()) + .map(str::to_string) + .or_else(|| { + env_lookup(REDUCTO_API_KEY_ENV) + .map(|key| key.trim().to_string()) + .filter(|key| !key.is_empty()) + }) + .ok_or(Error::MissingReductoApiKey)?; + Ok( + std::iter::once(("Authorization".into(), format!("Bearer {api_key}"))) + .chain(connection.extra_headers.clone()) + .collect(), + ) +} + +pub(super) async fn prepare_document( + client: &crate::ocr::OcrClient, + document: OcrDocument, + connection: &OcrConnection, + headers: &[(String, String)], +) -> Result { + if document.source().starts_with(REDUCTO_ID_PREFIX) { + if document.source()[REDUCTO_ID_PREFIX.len()..] + .trim() + .is_empty() + { + return Err(OcrRequestError::RequestField { + path: "document file id".into(), + } + .into()); + } + return Ok(document); + } + let inline = InlineDocument::parse(document.source())?.ok_or(OcrRequestError::ReductoSource)?; + let mime = inline.mime_type().to_string(); + let bytes = inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?; + let part = reqwest::multipart::Part::bytes(bytes) + .file_name("document") + .mime_str(&mime) + .map_err(|_| OcrRequestError::InvalidDataUri)?; + let builder = client + .provider_http() + .post(get_complete_url(connection.api_base.as_deref(), "upload")?) + .multipart(reqwest::multipart::Form::new().part("file", part)) + .timeout(connection.timeout); + let builder = crate::http_utils::with_headers( + builder, + headers, + crate::http_utils::HeaderPolicy::Except(&["content-type", "content-length"]), + ); + let response = crate::http_utils::http_request(builder) + .await + .map_err(crate::error::TransportError::from)?; + let uploaded = crate::ocr::client::read_json_response::< + crate::ocr::codecs::reducto::ReductoUploadResponse, + >(response, false) + .await? + .data; + let file_id = uploaded + .file_id + .as_deref() + .map(str::trim) + .filter(|id| !id.is_empty()); + let Some(file_id) = file_id else { + return Err(OcrResponseError::ResponseField { + path: "file_id".into(), + } + .into()); + }; + Ok(document.with_source(file_id.to_string())) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn explicit_key_precedes_environment_key() { + let connection = OcrConnection { + api_key: Some("passed-key".into()), + ..Default::default() + }; + let headers = validate_environment(&connection, &|_| Some("env-key".into())).unwrap(); + assert_eq!(headers[0].1, "Bearer passed-key"); + } + + #[test] + fn blank_explicit_key_uses_environment_key() { + let connection = OcrConnection { + api_key: Some(" ".into()), + ..Default::default() + }; + let headers = validate_environment(&connection, &|_| Some(" env-key ".into())).unwrap(); + assert_eq!(headers[0].1, "Bearer env-key"); + } + + #[test] + fn existing_authorization_skips_key_lookup() { + let connection = OcrConnection { + extra_headers: vec![("authorization".into(), "Bearer existing".into())], + ..Default::default() + }; + assert_eq!( + validate_environment(&connection, &|_| None).unwrap(), + connection.extra_headers + ); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/reducto/v3.rs b/litellm-rust/crates/core/src/ocr/adapters/reducto/v3.rs new file mode 100644 index 00000000000..a49f8105e26 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/reducto/v3.rs @@ -0,0 +1,45 @@ +use super::super::OcrAdapter; +use crate::ocr::OcrClient; +use crate::ocr::codecs::reducto::{self, ReductoResponse, ReductoV3Params}; +use crate::ocr::error::{OcrError, OcrResponseError}; +use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, build_http_request, credential_env, + guardrail_document, merge_extra_params, +}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; + +#[derive(Clone, Debug)] +pub(crate) struct ReductoV3Adapter; + +impl OcrAdapter for ReductoV3Adapter { + type ProviderResponse = ReductoResponse; + const PROVIDER: OcrProvider = OcrProvider::Reducto; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + let ParsedProviderParams { + known: params, + extra_params, + } = _prepare_ocr_request::(request)?; + let headers = super::validate_environment(&request.connection, &credential_env)?; + let url = super::get_complete_url(request.connection.api_base.as_deref(), "parse")?; + let document = guardrail_document(request, &url).await?; + let document = + super::prepare_document(client, document, &request.connection, &headers).await?; + let body = reducto::transform_v3_ocr_request(&request.model, document, ¶ms)?; + let body = merge_extra_params(&body, extra_params)?; + build_http_request(client, request, &url, &headers, &body) + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + reducto::transform_ocr_response(&request.model, response) + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/vertex/deepseek.rs b/litellm-rust/crates/core/src/ocr/adapters/vertex/deepseek.rs new file mode 100644 index 00000000000..ef188f8b9ac --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/vertex/deepseek.rs @@ -0,0 +1,134 @@ +use super::super::OcrAdapter; +use super::validate_destination; +use crate::Error; +use crate::auth::vertex::{self, VertexConfig}; +use crate::ocr::OcrClient; +use crate::ocr::codecs::deepseek::{self, DeepSeekOcrParams, DeepSeekOcrResponse}; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, +}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; +use crate::url_utils::ApiUrl; +const DEFAULT_API_BASE: &str = "https://aiplatform.googleapis.com"; +const MODEL_NAMESPACE: &str = "deepseek-ai"; +const DEFAULT_LOCATION: &str = "us-central1"; + +#[derive(Clone, Debug)] +pub(crate) struct VertexDeepSeekAdapter; + +impl OcrAdapter for VertexDeepSeekAdapter { + type ProviderResponse = DeepSeekOcrResponse; + const PROVIDER: OcrProvider = OcrProvider::VertexAi; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + validate_destination(&request.connection)?; + let ParsedProviderParams { + known: params, + extra_params: _extra_params, + } = _prepare_ocr_request::(request)?; + let config = VertexConfig::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)?; + let authentication = client + .vertex_auth() + .validate_environment( + request.connection.extra_headers.clone(), + request.connection.api_key.as_deref(), + &config, + &credential_env, + ) + .await + .map_err(Error::from)?; + let location = vertex::get_vertex_ai_location(&config, &credential_env) + .unwrap_or_else(|| DEFAULT_LOCATION.to_string()); + let url = get_complete_url( + request.connection.api_base.as_deref(), + &authentication.project_id, + &location, + )?; + let document = request.document.clone(); + let body = + deepseek::transform_ocr_request(&provider_model(&request.model), document, ¶ms)?; + transform_request_body(client, request, &url, &authentication.headers, body, |_| { + Ok(()) + }) + .await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + deepseek::transform_ocr_response(&request.model, response) + } +} + +fn provider_model(model: &str) -> String { + if model.starts_with(&format!("{MODEL_NAMESPACE}/")) { + model.to_string() + } else { + format!("{MODEL_NAMESPACE}/{model}") + } +} + +fn get_complete_url( + api_base: Option<&str>, + project: &str, + location: &str, +) -> Result { + let base = api_base + .map(str::trim) + .filter(|base| !base.is_empty()) + .unwrap_or(DEFAULT_API_BASE); + ApiUrl::parse(base) + .and_then(|url| { + url.complete_path(&[ + "v1", + "projects", + project, + "locations", + location, + "endpoints", + "openapi", + "chat", + "completions", + ]) + }) + .map(|url| url.into_string()) + .map_err(|_| { + OcrRequestError::RequestField { + path: "api_base".into(), + } + .into() + }) +} + +#[cfg(test)] +mod tests { + use super::{get_complete_url, provider_model}; + + #[test] + fn adapter_owns_model_namespace_and_endpoint() { + assert_eq!( + provider_model("deepseek-ocr-maas"), + "deepseek-ai/deepseek-ocr-maas" + ); + assert_eq!( + provider_model("deepseek-ai/deepseek-ocr-maas"), + "deepseek-ai/deepseek-ocr-maas" + ); + assert_eq!( + get_complete_url(None, "proj-1", "europe-west4").unwrap(), + "https://aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/endpoints/openapi/chat/completions" + ); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/vertex/mistral.rs b/litellm-rust/crates/core/src/ocr/adapters/vertex/mistral.rs new file mode 100644 index 00000000000..f3335bf497c --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/vertex/mistral.rs @@ -0,0 +1,154 @@ +use super::super::OcrAdapter; +use super::validate_destination; +use crate::Error; +use crate::auth::vertex::{self, VertexConfig}; +use crate::ocr::OcrClient; +use crate::ocr::codecs::mistral::{self, MistralOcrParams, MistralOcrResponse}; +use crate::ocr::document::{inline_remote_document, validate_inline_document}; +use crate::ocr::error::{OcrError, OcrRequestError, OcrResponseError}; +use crate::ocr::prepare::{ + _prepare_ocr_request, ParsedProviderParams, credential_env, transform_request_body, +}; +use crate::ocr::registry::OcrProvider; +use crate::ocr::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; +use crate::url_utils::ApiUrl; +const DEFAULT_LOCATION: &str = "us-central1"; + +#[derive(Clone, Debug)] +pub(crate) struct VertexMistralAdapter; + +impl OcrAdapter for VertexMistralAdapter { + type ProviderResponse = MistralOcrResponse; + const PROVIDER: OcrProvider = OcrProvider::VertexAi; + + async fn prepare_request( + &self, + request: &LiteLLMOcrRequest, + client: &OcrClient, + ) -> Result { + validate_destination(&request.connection)?; + let ParsedProviderParams { + known: params, + extra_params: _extra_params, + } = _prepare_ocr_request::(request)?; + let config = VertexConfig::from_sourced_optional_params( + &request.optional_params, + &request.input_sources, + ) + .map_err(Error::from)?; + let authentication = client + .vertex_auth() + .validate_environment( + request.connection.extra_headers.clone(), + request.connection.api_key.as_deref(), + &config, + &credential_env, + ) + .await + .map_err(Error::from)?; + let location = vertex::get_vertex_ai_location(&config, &credential_env) + .unwrap_or_else(|| DEFAULT_LOCATION.to_string()); + let url = get_complete_url( + request.connection.api_base.as_deref(), + &authentication.project_id, + &location, + &request.model, + )?; + let document = inline_remote_document( + client.document_fetcher(), + request.document.clone(), + &request.connection, + ) + .await?; + let body = mistral::transform_ocr_request(&request.model, document, ¶ms)?; + transform_request_body( + client, + request, + &url, + &authentication.headers, + body, + |body| validate_inline_document(&body.document), + ) + .await + } + + fn transform_ocr_response( + &self, + request: &LiteLLMOcrRequest, + response: Self::ProviderResponse, + ) -> Result { + mistral::transform_ocr_response(&request.model, response) + } +} + +fn get_complete_url( + api_base: Option<&str>, + project: &str, + location: &str, + model: &str, +) -> Result { + validate_location(location)?; + let default_base = format!("https://{location}-aiplatform.googleapis.com"); + let base = api_base + .map(str::trim) + .filter(|base| !base.is_empty()) + .unwrap_or(&default_base); + let prediction = format!("{model}:rawPredict"); + ApiUrl::parse(base) + .and_then(|url| { + url.complete_path(&[ + "v1", + "projects", + project, + "locations", + location, + "publishers", + "mistralai", + "models", + &prediction, + ]) + }) + .map(|url| url.into_string()) + .map_err(|_| { + OcrRequestError::RequestField { + path: "api_base".into(), + } + .into() + }) +} + +fn validate_location(location: &str) -> Result<(), OcrError> { + let valid = !location.is_empty() + && location + .bytes() + .all(|value| value.is_ascii_lowercase() || value.is_ascii_digit() || value == b'-') + && location + .as_bytes() + .first() + .is_some_and(u8::is_ascii_alphanumeric) + && location + .as_bytes() + .last() + .is_some_and(u8::is_ascii_alphanumeric); + if valid { + return Ok(()); + } + Err(OcrRequestError::RequestField { + path: "vertex_location".into(), + } + .into()) +} + +#[cfg(test)] +mod tests { + use super::get_complete_url; + + #[test] + fn endpoint_uses_location_project_and_model() { + assert_eq!( + get_complete_url(None, "proj-1", "europe-west4", "mistral-ocr-maas").unwrap(), + "https://europe-west4-aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict" + ); + assert!(get_complete_url(None, "proj-1", "attacker.example/path", "model").is_err()); + } +} diff --git a/litellm-rust/crates/core/src/ocr/adapters/vertex/mod.rs b/litellm-rust/crates/core/src/ocr/adapters/vertex/mod.rs new file mode 100644 index 00000000000..270c41e647d --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/adapters/vertex/mod.rs @@ -0,0 +1,21 @@ +mod deepseek; +mod mistral; + +use crate::Error; +use crate::auth::InputSource; +use crate::auth::error::AuthConfigurationError; +use crate::ocr::error::OcrError; +use crate::ocr::types::OcrConnection; + +pub(crate) use deepseek::VertexDeepSeekAdapter; +pub(crate) use mistral::VertexMistralAdapter; + +fn validate_destination(connection: &OcrConnection) -> Result<(), OcrError> { + if connection.api_base.is_some() && connection.api_base_source == InputSource::Request { + return Err(Error::from(crate::AuthError::Configuration( + AuthConfigurationError::RequestVertexCredentialDestination, + )) + .into()); + } + Ok(()) +} diff --git a/litellm-rust/crates/core/src/ocr/client.rs b/litellm-rust/crates/core/src/ocr/client.rs index 36bfc036bae..ab2d098d0bb 100644 --- a/litellm-rust/crates/core/src/ocr/client.rs +++ b/litellm-rust/crates/core/src/ocr/client.rs @@ -8,17 +8,28 @@ use super::handler::perform_ocr_request; use super::types::{LiteLLMOcrRequest, LiteLLMOcrResponse}; use super::wire::{DecodedOcrResponse, decode_response}; use crate::Error; +use crate::auth::vertex::VertexAuth; use crate::constants::OCR_CONNECT_TIMEOUT_SECS; use crate::error::TransportError; +use crate::media::MediaFetcher; #[derive(Clone)] pub struct OcrClient { provider_http: reqwest::Client, + polling_http: reqwest::Client, + document_fetcher: MediaFetcher, + vertex_auth: VertexAuth, } impl OcrClient { pub fn new(provider_http: reqwest::Client) -> Result { - Ok(Self { provider_http }) + let document_fetcher = MediaFetcher::new().map_err(TransportError::from)?; + Ok(Self { + provider_http, + polling_http: no_redirect_http()?, + document_fetcher, + vertex_auth: VertexAuth::default(), + }) } #[tracing::instrument( @@ -35,10 +46,35 @@ impl OcrClient { &self.provider_http } - #[cfg(test)] - pub(crate) fn for_test(provider_http: reqwest::Client) -> Self { - Self { provider_http } + pub(crate) fn polling_http(&self) -> &reqwest::Client { + &self.polling_http } + + pub(crate) fn document_fetcher(&self) -> &MediaFetcher { + &self.document_fetcher + } + + pub(crate) fn vertex_auth(&self) -> &VertexAuth { + &self.vertex_auth + } + + #[cfg(test)] + pub(crate) fn for_test(provider_http: reqwest::Client, document_http: reqwest::Client) -> Self { + Self { + provider_http, + polling_http: no_redirect_http().expect("test polling client builds"), + document_fetcher: MediaFetcher::for_test(document_http), + vertex_auth: VertexAuth::default(), + } + } +} + +fn no_redirect_http() -> Result { + reqwest::Client::builder() + .connect_timeout(Duration::from_secs(OCR_CONNECT_TIMEOUT_SECS)) + .redirect(reqwest::redirect::Policy::none()) + .build() + .map_err(TransportError::from) } pub async fn ocr(request: LiteLLMOcrRequest) -> Result { diff --git a/litellm-rust/crates/core/src/ocr/codecs/deepseek/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/deepseek/mod.rs new file mode 100644 index 00000000000..682b3addde7 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/deepseek/mod.rs @@ -0,0 +1,5 @@ +mod transformation; +mod types; + +pub(crate) use transformation::{transform_ocr_request, transform_ocr_response}; +pub(crate) use types::{DeepSeekOcrParams, DeepSeekOcrResponse}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs new file mode 100644 index 00000000000..98cfc0db78d --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/deepseek/transformation.rs @@ -0,0 +1,98 @@ +use serde::de::IntoDeserializer; +use serde_json::{Value, json}; + +use super::types::*; +use crate::ocr::error::{OcrRequestError, OcrResponseError}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub(crate) fn transform_ocr_request( + provider_model: &str, + document: OcrDocument, + params: &DeepSeekOcrParams, +) -> Result { + if document.source().is_empty() { + return Err(OcrRequestError::MissingField("document URL")); + } + Ok(DeepSeekOcrRequest { + model: provider_model.to_string(), + messages: vec![DeepSeekOcrMessage { + role: UserRole::User, + content: vec![document], + }], + params: params.clone(), + }) +} + +pub(crate) fn transform_ocr_response( + model: &str, + response: DeepSeekOcrResponse, +) -> Result { + let content = response + .choices + .into_iter() + .next() + .and_then(|choice| choice.message.content) + .ok_or(OcrResponseError::EmptyContent)?; + let decoded = decode_content(content)?; + let pages = match decoded.result.pages { + Some(pages) if !pages.is_empty() => pages + .into_iter() + .map(|page| serde_json::to_value(page).expect("DeepSeek page serializes")) + .collect(), + _ => vec![json!({ + "index":0, + "markdown":decoded.fallback_markdown, + "images":null + })], + }; + Ok(LiteLLMOcrResponse { + pages, + model: decoded.result.model.unwrap_or_else(|| model.to_string()), + document_annotation: decoded.result.document_annotation, + usage_info: decoded.result.usage_info.or(response.usage), + object: "ocr".into(), + extra_fields: decoded.result.extra_fields, + provider_native_response: None, + }) +} + +struct DecodedContent { + result: DeepSeekOcrResult, + fallback_markdown: String, +} + +fn decode_content(content: DeepSeekContent) -> Result { + let (result, fallback_markdown) = match content { + DeepSeekContent::Text(text) if text.is_empty() => { + return Err(OcrResponseError::EmptyContent); + } + DeepSeekContent::Text(text) => (decode_json_content(&text)?, text), + DeepSeekContent::Object(object) => { + let fallback = + serde_json::to_string(&object).map_err(|_| OcrResponseError::ResponseField { + path: "choices[0].message.content".into(), + })?; + (Some(object), fallback) + } + }; + Ok(DecodedContent { + result: result.unwrap_or_default(), + fallback_markdown, + }) +} + +fn decode_json_content(text: &str) -> Result, OcrResponseError> { + if !text.trim_start().starts_with('{') { + return Ok(None); + } + let value = match serde_json::from_str::(text) { + Ok(value) => value, + Err(_) => return Ok(None), + }; + serde_path_to_error::deserialize(value.into_deserializer()) + .map(Some) + .map_err(|error| OcrResponseError::ResponseField { + path: format!("choices[0].message.content.{}", error.path()), + }) +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/deepseek/types.rs b/litellm-rust/crates/core/src/ocr/codecs/deepseek/types.rs new file mode 100644 index 00000000000..0ce2d9913f7 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/deepseek/types.rs @@ -0,0 +1,95 @@ +use serde::{Deserialize, Serialize}; +use serde_json::{Map, Value}; + +#[derive(Clone, Debug, Default, Serialize, Deserialize)] +pub(crate) struct DeepSeekOcrParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub stream: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub top_p: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub n: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub stop: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +#[serde(untagged)] +pub(crate) enum StopSequences { + One(String), + Many(Vec), +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct DeepSeekOcrRequest { + pub model: String, + pub messages: Vec, + #[serde(flatten)] + pub params: DeepSeekOcrParams, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct DeepSeekOcrMessage { + pub role: UserRole, + pub content: Vec, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub(crate) enum UserRole { + User, +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct DeepSeekOcrResponse { + #[serde(default)] + pub choices: Vec, + pub usage: Option, +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct DeepSeekChoice { + pub message: DeepSeekResponseMessage, +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct DeepSeekResponseMessage { + pub content: Option, +} + +#[derive(Clone, Debug, Deserialize)] +#[serde(untagged)] +pub(crate) enum DeepSeekContent { + Text(String), + Object(DeepSeekOcrResult), +} + +#[derive(Clone, Debug, Default, Serialize, Deserialize)] +pub(crate) struct DeepSeekOcrResult { + #[serde(skip_serializing_if = "Option::is_none")] + pub pages: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub model: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub usage_info: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub document_annotation: Option, + #[serde(flatten)] + pub extra_fields: Map, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct DeepSeekPage { + #[serde(default)] + pub index: i64, + #[serde(default)] + pub markdown: String, + pub images: Option, + pub dimensions: Option, + #[serde(flatten)] + pub extra_fields: Map, +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/mod.rs new file mode 100644 index 00000000000..8031f2124a3 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/mod.rs @@ -0,0 +1,9 @@ +mod params; +mod transformation; +mod types; + +pub(crate) use params::{decode_input_params, map_ocr_params}; +pub(crate) use transformation::{transform_ocr_request, transform_ocr_response}; +pub(crate) use types::{ + AzureDocumentIntelligenceOperation, DocumentIntelligenceParams, OperationStatus, +}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/params.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/params.rs new file mode 100644 index 00000000000..85d1dafa542 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/params.rs @@ -0,0 +1,195 @@ +use std::collections::BTreeSet; + +use serde_json::{Map, Value}; + +use super::types::{ + DocumentIntelligenceInputParams, DocumentIntelligenceParams, FeaturesInput, PagesInput, +}; +use crate::ocr::error::OcrRequestError; +use crate::ocr::prepare::ParsedProviderParams; + +pub(crate) fn decode_input_params( + params: Map, + prefix: &str, +) -> Result, OcrRequestError> { + if let Some(Value::Array(pages)) = params.get("pages") { + if pages.iter().any(Value::is_boolean) { + return Err(OcrRequestError::Pages("boolean page index".into())); + } + if pages + .iter() + .any(|page| page.is_number() && page.as_i64().is_none()) + { + return Err(OcrRequestError::Pages("page index is out of range".into())); + } + if !pages.iter().all(Value::is_i64) && !pages.iter().all(Value::is_string) { + return Err(OcrRequestError::Pages("mixed page element types".into())); + } + } + crate::ocr::wire::decode_request_value(Value::Object(params), prefix) +} + +pub(crate) fn map_ocr_params( + params: DocumentIntelligenceInputParams, +) -> Result { + Ok(DocumentIntelligenceParams { + pages: params.pages.map(normalize_pages).transpose()?.flatten(), + features: params + .features + .map(normalize_features) + .transpose()? + .flatten(), + }) +} + +fn normalize_pages(pages: PagesInput) -> Result, OcrRequestError> { + let normalized = match pages { + PagesInput::ZeroBasedIndices(indices) => { + if indices.is_empty() { + return Ok(None); + } + indices + .into_iter() + .map(|page| { + if page < 0 { + return Err(OcrRequestError::Pages("negative page index".into())); + } + page.checked_add(1) + .ok_or_else(|| OcrRequestError::Pages("page index is out of range".into())) + }) + .collect::, _>>()? + .into_iter() + .map(|page| page.to_string()) + .collect::>() + .join(",") + } + PagesInput::NativeTokens(tokens) => { + if tokens.is_empty() { + return Ok(None); + } + tokens + .iter() + .map(|token| token.trim()) + .collect::>() + .join(",") + } + PagesInput::NativeRange(range) => range + .split(',') + .map(str::trim) + .collect::>() + .join(","), + }; + if !normalized.split(',').all(valid_page_token) { + return Err(OcrRequestError::Pages("invalid native page range".into())); + } + Ok(Some(normalized)) +} + +fn valid_page_token(token: &str) -> bool { + let mut parts = token.split('-'); + let start = parts.next().unwrap_or_default(); + if start.is_empty() || !start.chars().all(|character| character.is_ascii_digit()) { + return false; + } + match parts.next() { + None => true, + Some(end) => { + !end.is_empty() + && end.chars().all(|character| character.is_ascii_digit()) + && parts.next().is_none() + } + } +} + +fn normalize_features(features: FeaturesInput) -> Result, OcrRequestError> { + let tokens = match features { + FeaturesInput::Names(names) => names, + FeaturesInput::CommaSeparated(names) => names.split(',').map(str::to_string).collect(), + }; + if tokens.is_empty() { + return Ok(None); + } + let normalized = tokens.iter().map(|token| token.trim()).collect::>(); + if !normalized.iter().all(|token| { + let Some((first, rest)) = token.as_bytes().split_first() else { + return false; + }; + first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric) + }) { + return Err(OcrRequestError::Features); + } + Ok(Some(normalized.join(","))) +} + +#[cfg(test)] +mod tests { + use rstest::rstest; + use serde_json::{Value, json}; + + use super::*; + + fn map(value: Value) -> Result { + let fields = value.as_object().unwrap().clone(); + map_ocr_params(decode_input_params(fields, "optional_params")?.known) + } + + #[test] + fn input_params_retain_unknown_fields() { + let parsed = decode_input_params( + json!({ + "pages": [0], + "future_ocr_option": true, + "extra_body": {"provider_option": "value"} + }) + .as_object() + .unwrap() + .clone(), + "optional_params", + ) + .unwrap(); + + assert_eq!( + parsed.known.pages, + Some(PagesInput::ZeroBasedIndices(vec![0])) + ); + assert_eq!(parsed.extra_params["future_ocr_option"], true); + assert_eq!( + parsed.extra_params["extra_body"], + json!({"provider_option": "value"}) + ); + assert_eq!( + serde_json::to_value(map_ocr_params(parsed.known).unwrap()).unwrap(), + json!({"pages": "1", "features": null}) + ); + } + + #[rstest] + #[case(json!(["keyValuePairs"]), "keyValuePairs")] + #[case(json!(["keyValuePairs", "languages"]), "keyValuePairs,languages")] + #[case(json!("keyValuePairs"), "keyValuePairs")] + #[case(json!("keyValuePairs,languages"), "keyValuePairs,languages")] + #[case(json!("keyValuePairs, languages"), "keyValuePairs,languages")] + fn feature_mapping_matches_python(#[case] input: Value, #[case] expected: &str) { + assert_eq!( + map(json!({"features": input})).unwrap().features.as_deref(), + Some(expected) + ); + } + + #[rstest] + #[case(json!("keyValuePairs&pages=9"))] + #[case(json!("key value pairs"))] + #[case(json!(""))] + #[case(json!([1, 2]))] + #[case(json!([["keyValuePairs"]]))] + #[case(json!({"feature":"keyValuePairs"}))] + #[case(json!(5))] + fn invalid_feature_mapping_matches_python(#[case] input: Value) { + assert!(map(json!({"features": input})).is_err()); + } + + #[test] + fn empty_feature_list_is_omitted() { + assert_eq!(map(json!({"features": []})).unwrap().features, None); + } +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs new file mode 100644 index 00000000000..2b848fcfb7a --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/transformation.rs @@ -0,0 +1,111 @@ +use base64::{Engine, engine::general_purpose::STANDARD}; +use serde_json::{Map, Value, json}; + +use super::types::*; +use crate::constants::{AZURE_DI_DEFAULT_DPI, AZURE_DI_DEFAULT_HEIGHT, AZURE_DI_DEFAULT_WIDTH}; +use crate::ocr::document::InlineDocument; +use crate::ocr::error::{OcrRequestError, OcrResponseError}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; + +#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] +pub(crate) fn transform_ocr_request( + document: OcrDocument, +) -> Result { + let source = document.source(); + if source.is_empty() { + return Err(OcrRequestError::MissingField("document URL")); + } + Ok(if let Some(document) = InlineDocument::parse(source)? { + DocumentIntelligenceRequest::Base64Source( + STANDARD.encode(document.decode(crate::constants::OCR_INLINE_MAX_BYTES)?), + ) + } else { + DocumentIntelligenceRequest::UrlSource(source.to_string()) + }) +} + +pub(crate) fn transform_ocr_response( + model: &str, + response: AzureDocumentIntelligenceOperation, +) -> Result { + if response.status != Some(OperationStatus::Succeeded) { + return Err(OcrResponseError::OperationStatus( + response + .status + .map(|status| status.to_string()) + .unwrap_or_else(|| "None".into()), + )); + } + let result = response.analyze_result.unwrap_or_default(); + let pages = result + .pages + .into_iter() + .map(normalize_page) + .collect::, _>>()?; + let pages_processed = pages.len(); + let mut extra_fields = Map::new(); + extra_fields.insert("content".into(), option_value(result.content)); + extra_fields.insert("tables".into(), option_value(result.tables)); + extra_fields.insert( + "key_value_pairs".into(), + option_value(result.key_value_pairs), + ); + Ok(LiteLLMOcrResponse { + pages, + model: model.into(), + document_annotation: None, + usage_info: Some(json!({"pages_processed":pages_processed})), + object: "ocr".into(), + extra_fields, + provider_native_response: None, + }) +} + +fn normalize_page(page: AzureDocumentIntelligencePage) -> Result { + let index = page + .page_number + .unwrap_or(1) + .checked_sub(1) + .ok_or(OcrResponseError::NumericRange("page.pageNumber"))?; + let scale = if page.unit.as_deref().unwrap_or("inch") == "inch" { + AZURE_DI_DEFAULT_DPI as f64 + } else { + 1.0 + }; + let width = pixel_dimension( + page.width.unwrap_or(AZURE_DI_DEFAULT_WIDTH), + scale, + "page.width", + )?; + let height = pixel_dimension( + page.height.unwrap_or(AZURE_DI_DEFAULT_HEIGHT), + scale, + "page.height", + )?; + let markdown = page + .lines + .iter() + .map(|line| line.content.as_deref().unwrap_or_default()) + .collect::>() + .join("\n"); + Ok(json!({ + "index":index, + "markdown":markdown, + "images":null, + "dimensions":{"width":width,"height":height,"dpi":AZURE_DI_DEFAULT_DPI} + })) +} + +fn pixel_dimension(value: f64, scale: f64, field: &'static str) -> Result { + let value = value * scale; + if !value.is_finite() || value < i64::MIN as f64 || value > i64::MAX as f64 { + return Err(OcrResponseError::NumericRange(field)); + } + Ok(value.trunc() as i64) +} + +fn option_value(value: Option) -> Value { + value + .and_then(|value| serde_json::to_value(value).ok()) + .unwrap_or(Value::Null) +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/types.rs b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/types.rs new file mode 100644 index 00000000000..793f4547e99 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/document_intelligence/types.rs @@ -0,0 +1,138 @@ +use serde::{Deserialize, Deserializer, Serialize}; +use serde_json::{Map, Value}; + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub(crate) enum PagesInput { + ZeroBasedIndices(Vec), + NativeTokens(Vec), + NativeRange(String), +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +#[serde(untagged)] +pub(crate) enum FeaturesInput { + Names(Vec), + CommaSeparated(String), +} + +#[derive(Clone, Debug, Default, PartialEq, Serialize, Deserialize)] +pub(crate) struct DocumentIntelligenceInputParams { + pub pages: Option, + pub features: Option, +} + +#[derive(Clone, Debug, PartialEq, Serialize)] +pub(crate) struct DocumentIntelligenceParams { + pub pages: Option, + pub features: Option, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) enum DocumentIntelligenceRequest { + #[serde(rename = "urlSource")] + UrlSource(String), + #[serde(rename = "base64Source")] + Base64Source(String), +} + +#[derive(Clone, Debug, PartialEq)] +pub(crate) enum OperationStatus { + Succeeded, + Running, + NotStarted, + Failed, + Unknown(String), +} + +impl<'de> Deserialize<'de> for OperationStatus { + fn deserialize>(deserializer: D) -> Result { + Ok(match String::deserialize(deserializer)?.as_str() { + "succeeded" => Self::Succeeded, + "running" => Self::Running, + "notStarted" => Self::NotStarted, + "failed" => Self::Failed, + value => Self::Unknown(value.to_string()), + }) + } +} + +impl std::fmt::Display for OperationStatus { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(match self { + Self::Succeeded => "succeeded", + Self::Running => "running", + Self::NotStarted => "notStarted", + Self::Failed => "failed", + Self::Unknown(value) => value, + }) + } +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct AzureDocumentIntelligenceOperation { + pub status: Option, + #[serde(rename = "analyzeResult")] + pub analyze_result: Option, +} + +#[derive(Clone, Debug, Default, Deserialize)] +pub(crate) struct AzureDocumentIntelligenceAnalyzeResult { + pub content: Option, + #[serde(default)] + pub pages: Vec, + pub tables: Option>>, + #[serde(rename = "keyValuePairs")] + pub key_value_pairs: Option>>, +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct AzureDocumentIntelligencePage { + #[serde(rename = "pageNumber", default, deserialize_with = "optional_i64")] + pub page_number: Option, + #[serde(default, deserialize_with = "optional_f64")] + pub width: Option, + #[serde(default, deserialize_with = "optional_f64")] + pub height: Option, + pub unit: Option, + #[serde(default)] + pub lines: Vec, +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct AzureDocumentIntelligenceLine { + pub content: Option, +} + +fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::deserialize(deserializer)? { + None | Some(Value::Null) => Ok(None), + Some(Value::Number(number)) => number + .as_i64() + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected an integer")), + Some(Value::String(value)) => value + .parse::() + .map(Some) + .map_err(|_| serde::de::Error::custom("expected an integer")), + Some(_) => Err(serde::de::Error::custom("expected an integer")), + } +} + +fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::deserialize(deserializer)? { + None | Some(Value::Null) => Ok(None), + Some(Value::Number(number)) => number + .as_f64() + .filter(|value| value.is_finite()) + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected a finite number")), + Some(Value::String(value)) => value + .parse::() + .ok() + .filter(|value| value.is_finite()) + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected a finite number")), + Some(_) => Err(serde::de::Error::custom("expected a number")), + } +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs index cd0a1dc6b17..5bd7e555a1e 100644 --- a/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs +++ b/litellm-rust/crates/core/src/ocr/codecs/mistral/transformation.rs @@ -36,6 +36,101 @@ mod tests { use rstest::rstest; use serde_json::{Value, json}; + fn mapped_params(value: Value) -> Value { + serde_json::to_value(serde_json::from_value::(value).unwrap()).unwrap() + } + + fn document() -> OcrDocument { + serde_json::from_value( + json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), + ) + .unwrap() + } + + #[rstest] + fn extract_header_is_a_supported_ocr_param() { + assert_eq!( + mapped_params(json!({"extract_header":true}))["extract_header"], + true + ); + } + + #[rstest] + fn extract_footer_is_a_supported_ocr_param() { + assert_eq!( + mapped_params(json!({"extract_footer":false}))["extract_footer"], + false + ); + } + + #[rstest] + fn existing_ocr_params_remain_supported() { + let mapped = mapped_params(json!({ + "pages":[0,2], + "include_image_base64":true, + "image_limit":2, + "image_min_size":100, + "bbox_annotation_format":{"type":"json_schema"}, + "document_annotation_format":{"type":"json_schema"} + })); + assert_eq!(mapped["pages"], json!([0, 2])); + assert_eq!(mapped["include_image_base64"], true); + assert_eq!(mapped["image_limit"], 2); + assert_eq!(mapped["image_min_size"], 100); + assert_eq!(mapped["bbox_annotation_format"]["type"], "json_schema"); + assert_eq!(mapped["document_annotation_format"]["type"], "json_schema"); + } + + #[rstest] + fn map_ocr_params_forwards_extract_header() { + assert_eq!( + mapped_params(json!({"extract_header":true}))["extract_header"], + true + ); + } + + #[rstest] + fn map_ocr_params_forwards_extract_footer() { + assert_eq!( + mapped_params(json!({"extract_footer":true}))["extract_footer"], + true + ); + } + + #[rstest] + fn map_ocr_params_forwards_extract_header_and_footer() { + let mapped = mapped_params(json!({"extract_header":true,"extract_footer":false})); + assert_eq!(mapped["extract_header"], true); + assert_eq!(mapped["extract_footer"], false); + } + + #[rstest] + fn map_ocr_params_drops_unknown_params() { + let mapped = mapped_params(json!({"extract_header":true,"unsupported_param":"value"})); + assert_eq!(mapped["extract_header"], true); + assert!(mapped.get("unsupported_param").is_none()); + } + + #[rstest] + #[case("table_format", json!("html"))] + #[case("confidence_scores_granularity", json!("word"))] + #[case("document_annotation_prompt", json!("extract"))] + #[case("include_blocks", json!(true))] + #[case("id", json!("req-123"))] + fn new_ocr_params_are_supported(#[case] name: &str, #[case] value: Value) { + assert_eq!(mapped_params(json!({name:value.clone()}))[name], value); + } + + #[rstest] + #[case("table_format", json!("html"))] + #[case("confidence_scores_granularity", json!("word"))] + #[case("document_annotation_prompt", json!("extract"))] + #[case("include_blocks", json!(true))] + #[case("id", json!("req-123"))] + fn map_ocr_params_forwards_new_ocr_params(#[case] name: &str, #[case] value: Value) { + assert_eq!(mapped_params(json!({name:value.clone()}))[name], value); + } + #[rstest] #[case("pages", json!([0, 2]))] #[case("include_image_base64", json!(true))] @@ -53,50 +148,81 @@ mod tests { fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) { let params: MistralOcrParams = serde_json::from_value(json!({name: value.clone()})).unwrap(); - let document: OcrDocument = serde_json::from_value( - json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), - ) - .unwrap(); let result = - serde_json::to_value(transform_ocr_request("model", document, ¶ms).unwrap()) + serde_json::to_value(transform_ocr_request("model", document(), ¶ms).unwrap()) .unwrap(); assert_eq!(result["model"], "model"); assert_eq!(result[name], value); } - #[test] - fn request_mapping_filters_unknown_fields() { - let params: MistralOcrParams = serde_json::from_value(json!({"unknown": true})).unwrap(); - let document: OcrDocument = serde_json::from_value( - json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), + #[rstest] + #[case("table_format", json!("html"))] + #[case("confidence_scores_granularity", json!("word"))] + #[case("document_annotation_prompt", json!("extract"))] + #[case("id", json!("req-123"))] + #[case("extract_header", json!(true))] + #[case("include_blocks", json!(true))] + #[case("pages", json!([0,1]))] + fn transform_ocr_request_includes_each_optional_param( + #[case] name: &str, + #[case] value: Value, + ) { + let params: MistralOcrParams = serde_json::from_value(json!({name:value.clone()})).unwrap(); + let result = serde_json::to_value( + transform_ocr_request("mistral-ocr-latest", document(), ¶ms).unwrap(), ) .unwrap(); - let result = - serde_json::to_value(transform_ocr_request("model", document, ¶ms).unwrap()) - .unwrap(); - assert!(result.get("unknown").is_none()); + assert_eq!(result[name], value); + assert_eq!(result["model"], "mistral-ocr-latest"); } - #[test] - fn response_preserves_provider_fields() { + #[rstest] + fn transform_ocr_request_includes_multiple_new_params() { + let params: MistralOcrParams = serde_json::from_value(json!({ + "table_format":"html", + "confidence_scores_granularity":"page", + "extract_header":true + })) + .unwrap(); + let result = serde_json::to_value( + transform_ocr_request("mistral-ocr-latest", document(), ¶ms).unwrap(), + ) + .unwrap(); + assert_eq!(result["table_format"], "html"); + assert_eq!(result["confidence_scores_granularity"], "page"); + assert_eq!(result["extract_header"], true); + } + + #[rstest] + fn transform_ocr_response_preserves_blocks_and_confidence_scores() { let response: MistralOcrResponse = serde_json::from_value(json!({ - "pages":[{"index":0,"markdown":"hello","header":"head","confidence_scores":{"mean":0.99}}], + "pages":[{"index":0,"markdown":"hello","blocks":[{"type":"title"}],"confidence_scores":{"mean":0.99}}], "model":"returned-model", - "usage_info":{"pages_processed":1,"future_counter":5}, - "future_response_field":"kept" + "usage_info":{"pages_processed":1} })) .unwrap(); let result = transform_ocr_response("model", response) .unwrap() .into_json(); - assert_eq!(result["pages"][0]["header"], "head"); - assert_eq!(result["usage_info"]["future_counter"], 5); - assert_eq!(result["future_response_field"], "kept"); - assert_eq!(result["model"], "returned-model"); + assert_eq!(result["pages"][0]["blocks"][0]["type"], "title"); + assert_eq!(result["pages"][0]["confidence_scores"]["mean"], 0.99); } - #[test] - fn response_rejects_null_pages() { - assert!(serde_json::from_value::(json!({"pages":null})).is_err()); + #[rstest] + fn transform_ocr_response_preserves_ocr4_page_fields() { + let page = json!({ + "index":0, + "markdown":"table page", + "tables":[{"rows":2,"cols":3}], + "hyperlinks":["https://example.com"], + "header":"header", + "footer":"footer" + }); + let response: MistralOcrResponse = + serde_json::from_value(json!({"pages":[page.clone()]})).unwrap(); + let result = transform_ocr_response("model", response) + .unwrap() + .into_json(); + assert_eq!(result["pages"][0], page); } } diff --git a/litellm-rust/crates/core/src/ocr/codecs/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/mod.rs index 170ef5f68a7..7c752749901 100644 --- a/litellm-rust/crates/core/src/ocr/codecs/mod.rs +++ b/litellm-rust/crates/core/src/ocr/codecs/mod.rs @@ -1 +1,4 @@ +pub(crate) mod deepseek; +pub(crate) mod document_intelligence; pub(crate) mod mistral; +pub(crate) mod reducto; diff --git a/litellm-rust/crates/core/src/ocr/codecs/reducto/mod.rs b/litellm-rust/crates/core/src/ocr/codecs/reducto/mod.rs new file mode 100644 index 00000000000..3fff40451c6 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/reducto/mod.rs @@ -0,0 +1,9 @@ +mod transformation; +mod types; + +pub(crate) use transformation::{ + transform_legacy_ocr_request, transform_ocr_response, transform_v3_ocr_request, +}; +pub(crate) use types::{ + ReductoLegacyParams, ReductoResponse, ReductoUploadResponse, ReductoV3Params, +}; diff --git a/litellm-rust/crates/core/src/ocr/codecs/reducto/transformation.rs b/litellm-rust/crates/core/src/ocr/codecs/reducto/transformation.rs new file mode 100644 index 00000000000..7073643f6b6 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/reducto/transformation.rs @@ -0,0 +1,115 @@ +use std::collections::BTreeMap; + +use serde_json::{Value, json}; + +use super::types::*; +use crate::ocr::error::{OcrRequestError, OcrResponseError}; +use crate::ocr::types::{LiteLLMOcrResponse, OcrDocument}; + +#[tracing::instrument( + name = "transform_ocr_request", + target = "litellm::function_trace", + level = "trace", + skip_all +)] +pub(crate) fn transform_v3_ocr_request( + _model: &str, + document: OcrDocument, + params: &ReductoV3Params, +) -> Result { + Ok(ReductoV3Request { + input: document.source().to_string(), + params: params.clone(), + }) +} + +#[tracing::instrument( + name = "transform_ocr_request", + target = "litellm::function_trace", + level = "trace", + skip_all +)] +pub(crate) fn transform_legacy_ocr_request( + _model: &str, + document: OcrDocument, + params: &ReductoLegacyParams, +) -> Result { + Ok(ReductoLegacyRequest { + document_url: document.source().to_string(), + options: params.enhance.as_ref().map(|_| params.clone()), + }) +} + +pub(crate) fn transform_ocr_response( + model: &str, + response: ReductoResponse, +) -> Result { + let result = match response.result { + Some(result) => result.unwrap_or_default(), + None => ReductoResult { + chunks: response.chunks, + }, + }; + let usage = response.usage.unwrap_or_default(); + Ok(LiteLLMOcrResponse { + pages: build_pages(result.chunks.unwrap_or_default()), + model: model.to_string(), + document_annotation: None, + usage_info: Some(json!({ + "pages_processed": usage.num_pages, + "credits": usage.credits, + })), + object: "ocr".to_string(), + extra_fields: serde_json::Map::new(), + provider_native_response: None, + }) +} + +fn build_pages(chunks: Vec) -> Vec { + let blocks_by_page = chunks + .iter() + .flat_map(|chunk| chunk.blocks.iter().flatten()) + .filter_map(|block| block.bbox.as_ref()?.page.map(|page| (page, block))) + .fold( + BTreeMap::>::new(), + |mut pages, (page, block)| { + pages.entry(page).or_default().push(block); + pages + }, + ); + if blocks_by_page.is_empty() { + let markdown = join_content(chunks.iter().map(|chunk| chunk.content.as_deref())); + return if markdown.is_empty() { + Vec::new() + } else { + vec![page(0, markdown, None)] + }; + } + blocks_by_page + .into_iter() + .map(|(index, blocks)| { + let markdown = join_content(blocks.iter().map(|block| block.content.as_deref())); + page( + index.saturating_sub(1).max(0), + markdown, + Some(json!(blocks)), + ) + }) + .collect() +} + +fn join_content<'a>(content: impl Iterator>) -> String { + content + .flatten() + .filter(|text| !text.is_empty()) + .collect::>() + .join("\n\n") +} + +fn page(index: i64, markdown: String, blocks: Option) -> Value { + let mut result = json!({"index":index,"markdown":markdown,"images":null}); + if let (Value::Object(fields), Some(blocks)) = (&mut result, blocks) { + fields.insert("blocks".into(), blocks); + } + result +} diff --git a/litellm-rust/crates/core/src/ocr/codecs/reducto/types.rs b/litellm-rust/crates/core/src/ocr/codecs/reducto/types.rs new file mode 100644 index 00000000000..c03720cc8ae --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/codecs/reducto/types.rs @@ -0,0 +1,128 @@ +use serde::{Deserialize, Deserializer, Serialize}; +use serde_json::{Map, Value}; + +#[derive(Clone, Debug, Default, Serialize, Deserialize)] +pub(crate) struct ReductoV3Params { + #[serde(skip_serializing_if = "Option::is_none")] + pub formatting: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub retrieval: Option>, + #[serde(skip_serializing_if = "Option::is_none")] + pub settings: Option>, +} + +#[derive(Clone, Debug, Default, Serialize, Deserialize)] +pub(crate) struct ReductoLegacyParams { + #[serde(skip_serializing_if = "Option::is_none")] + pub enhance: Option>, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct ReductoV3Request { + pub input: String, + #[serde(flatten)] + pub params: ReductoV3Params, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct ReductoLegacyRequest { + pub document_url: String, + #[serde(skip_serializing_if = "Option::is_none")] + pub options: Option, +} + +#[derive(Deserialize)] +pub(crate) struct ReductoUploadResponse { + pub file_id: Option, +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct ReductoResponse { + #[serde(default, deserialize_with = "present_nullable")] + pub result: Option>, + pub usage: Option, + #[serde(default)] + pub chunks: Option>, +} + +fn present_nullable<'de, D: Deserializer<'de>, T: Deserialize<'de>>( + deserializer: D, +) -> Result>, D::Error> { + Option::::deserialize(deserializer).map(Some) +} + +#[derive(Clone, Debug, Default, Deserialize)] +pub(crate) struct ReductoResult { + pub chunks: Option>, +} + +#[derive(Clone, Debug, Default, Deserialize)] +pub(crate) struct ReductoUsage { + #[serde(default, deserialize_with = "optional_i64")] + pub num_pages: Option, + #[serde(default, deserialize_with = "optional_f64")] + pub credits: Option, +} + +#[derive(Clone, Debug, Deserialize)] +pub(crate) struct ReductoChunk { + pub content: Option, + pub blocks: Option>, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct ReductoBlock { + #[serde(skip_serializing_if = "Option::is_none")] + pub content: Option, + #[serde(skip_serializing_if = "Option::is_none")] + pub bbox: Option, + #[serde(flatten)] + pub extra_fields: Map, +} + +#[derive(Clone, Debug, Serialize, Deserialize)] +pub(crate) struct ReductoBoundingBox { + #[serde(default, deserialize_with = "optional_i64")] + pub page: Option, + #[serde(flatten)] + pub extra_fields: Map, +} + +fn optional_i64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::deserialize(deserializer)? { + None | Some(Value::Null) => Ok(None), + Some(Value::Number(number)) => number + .as_i64() + .or_else(|| number.as_f64().and_then(checked_truncated_i64)) + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected an integer")), + Some(Value::String(value)) => value + .trim() + .parse::() + .map(Some) + .map_err(|_| serde::de::Error::custom("expected an integer")), + Some(Value::Bool(value)) => Ok(Some(i64::from(value))), + Some(_) => Ok(None), + } +} + +fn optional_f64<'de, D: Deserializer<'de>>(deserializer: D) -> Result, D::Error> { + match Option::::deserialize(deserializer)? { + None | Some(Value::Null) => Ok(None), + Some(Value::Number(number)) => number + .as_f64() + .map(Some) + .ok_or_else(|| serde::de::Error::custom("expected a number")), + Some(Value::String(value)) => value + .trim() + .parse::() + .map(Some) + .map_err(|_| serde::de::Error::custom("expected a number")), + Some(_) => Ok(None), + } +} + +fn checked_truncated_i64(value: f64) -> Option { + (value.is_finite() && value >= i64::MIN as f64 && value <= i64::MAX as f64) + .then(|| value.trunc() as i64) +} diff --git a/litellm-rust/crates/core/src/ocr/document.rs b/litellm-rust/crates/core/src/ocr/document.rs new file mode 100644 index 00000000000..e89b1c5c569 --- /dev/null +++ b/litellm-rust/crates/core/src/ocr/document.rs @@ -0,0 +1,212 @@ +use base64::{Engine, engine::general_purpose::STANDARD}; +use data_url::mime::Mime; +use data_url::{DataUrl, DataUrlError, forgiving_base64::DecodeError}; +use reqwest::Url; + +use super::error::{OcrError, OcrRequestError, OcrResponseError}; +use super::types::{OcrConnection, OcrDocument}; +use crate::constants::OCR_MAX_FETCH_REDIRECTS; +use crate::error::{MediaError, TransportError}; +use crate::media::{DownloadPolicy, MediaFetcher}; + +pub(crate) struct InlineDocument<'a>(DataUrl<'a>); + +impl<'a> InlineDocument<'a> { + pub(crate) fn parse(source: &'a str) -> Result, OcrRequestError> { + match DataUrl::process(source) { + Ok(url) => Ok(Some(Self(url))), + Err(DataUrlError::NotADataUrl) => Ok(None), + Err(DataUrlError::NoComma) => Err(OcrRequestError::InvalidDataUri), + } + } + + pub(crate) fn mime_type(&self) -> &Mime { + self.0.mime_type() + } + + pub(crate) fn decode(&self, max_bytes: usize) -> Result, OcrRequestError> { + let mut body = Vec::new(); + self.0 + .decode(|bytes| { + if bytes.len() > max_bytes.saturating_sub(body.len()) { + return Err(OcrRequestError::InlineDocumentTooLarge); + } + body.extend_from_slice(bytes); + Ok(()) + }) + .map_err(|error| match error { + DecodeError::InvalidBase64(_) => OcrRequestError::InvalidDataUri, + DecodeError::WriteError(error) => error, + })?; + Ok(body) + } +} + +pub(crate) fn validate_inline_document(document: &OcrDocument) -> Result<(), OcrRequestError> { + let inline = + InlineDocument::parse(document.source())?.ok_or(OcrRequestError::InvalidDataUri)?; + inline.decode(crate::constants::OCR_INLINE_MAX_BYTES)?; + Ok(()) +} + +pub(crate) async fn inline_remote_document( + fetcher: &MediaFetcher, + document: OcrDocument, + connection: &OcrConnection, +) -> Result { + let source = document.source(); + if !source.starts_with("http://") && !source.starts_with("https://") { + validate_inline_document(&document)?; + return Ok(document); + } + let url = Url::parse(source).map_err(|_| OcrRequestError::RequestField { + path: "document URL".into(), + })?; + let downloaded = fetcher + .fetch( + url, + DownloadPolicy { + timeout: connection.timeout, + max_bytes: connection.max_download_bytes, + max_redirects: OCR_MAX_FETCH_REDIRECTS, + }, + ) + .await + .map_err(map_media_error)?; + let result = document.with_source(format!( + "data:{};base64,{}", + downloaded.content_type, + STANDARD.encode(downloaded.bytes) + )); + validate_inline_document(&result)?; + Ok(result) +} + +fn map_media_error(error: MediaError) -> OcrError { + match error { + MediaError::BlockedUrl => OcrRequestError::BlockedDocumentUrl.into(), + MediaError::DownloadDisabled => OcrRequestError::DownloadDisabled.into(), + MediaError::DownloadTooLarge => OcrRequestError::DownloadTooLarge.into(), + MediaError::TooManyRedirects => OcrRequestError::TooManyRedirects.into(), + MediaError::MissingRedirectLocation => OcrResponseError::MissingRedirectLocation.into(), + MediaError::InvalidRedirect => OcrResponseError::InvalidRedirect.into(), + MediaError::Http(status) => TransportError::Http { + status, + body: "OCR document download failed".into(), + } + .into(), + MediaError::Timeout => { + TransportError::Network("OCR document download timed out".into()).into() + } + MediaError::Transport(error) => error.into(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use serde_json::Map; + + fn document(source: &str) -> OcrDocument { + OcrDocument::DocumentUrl { + document_url: source.into(), + extra_fields: Map::new(), + } + } + + #[test] + fn decodes_data_urls_and_limits_decoded_size() { + for (source, expected) in [ + ("data:application/pdf;base64,YWJj", b"abc".as_slice()), + ("DATA:application/pdf;BASE64,YWI", b"ab".as_slice()), + ("data:,a%20b%00%FF", b"a b\0\xff".as_slice()), + ] { + let inline = InlineDocument::parse(source).unwrap().unwrap(); + assert_eq!(inline.decode(expected.len()).unwrap(), expected); + assert_eq!( + inline.decode(expected.len() - 1), + Err(OcrRequestError::InlineDocumentTooLarge) + ); + } + } + + #[test] + fn preserves_mime_parameters_and_standard_default() { + let inline = InlineDocument::parse("data:application/pdf;version=1.7;base64,YQ==") + .unwrap() + .unwrap(); + assert!(inline.mime_type().matches("application", "pdf")); + assert_eq!(inline.mime_type().get_parameter("version"), Some("1.7")); + let default = InlineDocument::parse("data:,a").unwrap().unwrap(); + assert!(default.mime_type().matches("text", "plain")); + assert_eq!( + default.mime_type().get_parameter("charset"), + Some("US-ASCII") + ); + } + + #[test] + fn rejects_invalid_inline_documents() { + for source in [ + "https://example.com/document.pdf", + "data:application/pdf;base64", + "data:application/pdf;base64,INVALID!", + ] { + assert!(validate_inline_document(&document(source)).is_err()); + } + } + + #[tokio::test] + async fn remote_conversion_preserves_kind_and_isolates_provider_credentials() { + use tokio::io::{AsyncReadExt, AsyncWriteExt}; + use tokio::net::TcpListener; + + let listener = TcpListener::bind("127.0.0.1:0").await.unwrap(); + let address = listener.local_addr().unwrap(); + let server = tokio::spawn(async move { + let (mut socket, _) = listener.accept().await.unwrap(); + let mut request = vec![0_u8; 2048]; + let count = socket.read(&mut request).await.unwrap(); + socket + .write_all(b"HTTP/1.1 200 OK\r\nContent-Type: image/png; charset=binary\r\nContent-Length: 3\r\nConnection: close\r\n\r\nabc") + .await + .unwrap(); + String::from_utf8_lossy(&request[..count]).into_owned() + }); + let mut provider_headers = reqwest::header::HeaderMap::new(); + provider_headers.insert( + reqwest::header::AUTHORIZATION, + reqwest::header::HeaderValue::from_static("Bearer provider-secret"), + ); + let provider_http = reqwest::Client::builder() + .default_headers(provider_headers) + .build() + .unwrap(); + let document_http = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .unwrap(); + let client = super::super::OcrClient::for_test(provider_http, document_http); + let converted = inline_remote_document( + client.document_fetcher(), + OcrDocument::ImageUrl { + image_url: format!("http://{address}/image"), + extra_fields: Map::from_iter([("detail".into(), serde_json::json!("high"))]), + }, + &OcrConnection::default(), + ) + .await + .unwrap(); + let request = server.await.unwrap(); + + assert_eq!( + converted, + OcrDocument::ImageUrl { + image_url: "data:image/png;base64,YWJj".into(), + extra_fields: Map::from_iter([("detail".into(), serde_json::json!("high"))]), + } + ); + assert!(!request.to_ascii_lowercase().contains("authorization")); + assert!(!request.contains("provider-secret")); + } +} diff --git a/litellm-rust/crates/core/src/ocr/error.rs b/litellm-rust/crates/core/src/ocr/error.rs index 50793aad566..522d059ec48 100644 --- a/litellm-rust/crates/core/src/ocr/error.rs +++ b/litellm-rust/crates/core/src/ocr/error.rs @@ -10,12 +10,52 @@ pub enum OcrRequestError { RequestField { path: String }, #[error("missing required field: {0}")] MissingField(&'static str), + #[error("invalid OCR document data URI")] + InvalidDataUri, + #[error("Reducto requires a reducto:// id or a data URI")] + ReductoSource, + #[error("inline OCR document exceeds the size limit")] + InlineDocumentTooLarge, + #[error("OCR document URL is blocked by network policy")] + BlockedDocumentUrl, + #[error("OCR document downloads are disabled")] + DownloadDisabled, + #[error("OCR document download exceeds the size limit")] + DownloadTooLarge, + #[error("OCR document download exceeded the redirect limit")] + TooManyRedirects, + #[error("invalid OCR pages: {0}")] + Pages(String), + #[error("invalid OCR features")] + Features, + #[error("OCR model cannot be a dot segment")] + DotModel, } #[derive(Debug, Clone, PartialEq, Eq, Error)] pub enum OcrResponseError { #[error("invalid OCR response field: {path}")] ResponseField { path: String }, + #[error("OCR response is missing non-empty content")] + EmptyContent, + #[error("OCR document redirect is missing a location")] + MissingRedirectLocation, + #[error("OCR document redirect location is invalid")] + InvalidRedirect, + #[error("OCR operation ended with status {0}")] + OperationStatus(String), + #[error("OCR response numeric value is out of range: {0}")] + NumericRange(&'static str), +} + +#[derive(Debug, Clone, PartialEq, Eq, Error)] +pub enum OcrPollingError { + #[error("OCR accepted response is missing a valid operation-location")] + PollLocation, + #[error("OCR operation-location must use the submission origin without credentials")] + PollOrigin, + #[error("OCR polling timed out")] + PollTimeout, } #[derive(Debug, Error)] @@ -27,6 +67,8 @@ pub enum OcrError { #[error("{0}")] Transport(#[from] TransportError), #[error("{0}")] + Polling(#[from] OcrPollingError), + #[error("{0}")] Public(#[from] crate::Error), } @@ -36,6 +78,7 @@ impl From for crate::Error { OcrError::Request(error) => error.into(), OcrError::Response(error) => error.into(), OcrError::Transport(error) => error.into(), + OcrError::Polling(error) => crate::Error::InvalidResponse(error.to_string()), OcrError::Public(error) => error, } } diff --git a/litellm-rust/crates/core/src/ocr/mod.rs b/litellm-rust/crates/core/src/ocr/mod.rs index 78c86aeaa0e..1e975c3f521 100644 --- a/litellm-rust/crates/core/src/ocr/mod.rs +++ b/litellm-rust/crates/core/src/ocr/mod.rs @@ -1,21 +1,39 @@ mod adapters; pub mod client; mod codecs; +mod document; pub mod error; mod handler; pub mod hooks; mod prepare; mod registry; -pub mod transformation; pub mod types; pub mod wire; pub use client::{OcrClient, ocr}; pub use types::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrConnection, OcrDocument}; +#[cfg(test)] +#[path = "../../tests/azure_ai_ocr.rs"] +mod azure_ai_tests; +#[cfg(test)] +#[path = "../../tests/azure_document_intelligence_ocr.rs"] +mod azure_document_intelligence_tests; +#[cfg(test)] +#[path = "../../tests/deepseek_ocr.rs"] +mod deepseek_tests; +#[cfg(test)] +#[path = "../../tests/reducto_ocr.rs"] +mod reducto_tests; #[cfg(test)] #[path = "../../tests/ocr/support.rs"] pub(crate) mod test_support; #[cfg(test)] #[path = "../../tests/ocr.rs"] pub(crate) mod tests; +#[cfg(test)] +#[path = "../../tests/vertex_ai_deepseek_ocr.rs"] +mod vertex_ai_deepseek_tests; +#[cfg(test)] +#[path = "../../tests/vertex_ai_ocr.rs"] +mod vertex_ai_tests; diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index cff40b7de50..bf6f924088c 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -4,7 +4,7 @@ use serde_json::{Map, Value}; use super::OcrClient; use super::error::{OcrError, OcrRequestError}; use super::hooks::OcrDuringCallRequest; -use super::types::LiteLLMOcrRequest; +use super::types::{LiteLLMOcrRequest, OcrDocument}; #[derive(Debug, Deserialize)] pub(crate) struct ParsedProviderParams { @@ -24,6 +24,39 @@ pub(crate) fn _prepare_ocr_request( ) } +pub(crate) fn merge_extra_params( + body: &B, + extra_params: Map, +) -> Result { + let Value::Object(fields) = + serde_json::to_value(body).map_err(|_| OcrRequestError::RequestField { + path: "body".into(), + })? + else { + return Err(OcrRequestError::RequestField { + path: "body".into(), + }); + }; + let extra_body = extra_params + .get("extra_body") + .and_then(Value::as_object) + .cloned() + .unwrap_or_default() + .into_iter() + .collect::>(); + Ok(Value::Object( + fields + .into_iter() + .chain( + extra_params + .into_iter() + .filter(|(name, _)| name != "extra_body"), + ) + .chain(extra_body) + .collect(), + )) +} + pub(crate) async fn transform_request_body( client: &OcrClient, request: &LiteLLMOcrRequest, @@ -77,6 +110,29 @@ pub(crate) fn build_http_request( .map_err(OcrError::from) } +pub(crate) async fn guardrail_document( + request: &LiteLLMOcrRequest, + url: &str, +) -> Result { + if !request.hooks.has_guardrails() { + return Ok(request.document.clone()); + } + let changed = request + .hooks + .during_call(OcrDuringCallRequest { + model: request.model.clone(), + custom_llm_provider: request.adapter.provider().as_str().into(), + url: url.into(), + body: serde_json::to_value(&request.document).map_err(|_| { + OcrRequestError::RequestField { + path: "document".into(), + } + })?, + }) + .await?; + super::wire::decode_request_value(changed.body, "guardrail.document").map_err(OcrError::from) +} + #[derive(Serialize)] struct OcrWireBody { #[serde(flatten)] @@ -107,7 +163,6 @@ impl OcrWireBody { pub(crate) fn credential_env(name: &str) -> Option { std::env::var(name).ok() } - #[cfg(test)] mod tests { use serde_json::json; diff --git a/litellm-rust/crates/core/src/ocr/registry.rs b/litellm-rust/crates/core/src/ocr/registry.rs index 9bae8153ee8..1b20a91143b 100644 --- a/litellm-rust/crates/core/src/ocr/registry.rs +++ b/litellm-rust/crates/core/src/ocr/registry.rs @@ -24,12 +24,18 @@ super::adapters::for_each_ocr_adapter!(define_adapter_types); #[derive(Clone, Copy, Debug, PartialEq, Eq)] pub(crate) enum OcrProvider { Mistral, + AzureAi, + Reducto, + VertexAi, } impl OcrProvider { pub(crate) const fn as_str(self) -> &'static str { match self { Self::Mistral => "mistral", + Self::AzureAi => "azure_ai", + Self::Reducto => "reducto", + Self::VertexAi => "vertex_ai", } } } @@ -45,9 +51,78 @@ pub(crate) fn resolve_wire_adapter( }); let typed_provider = match provider.custom_llm_provider { "mistral" => OcrProvider::Mistral, + "azure_ai" => OcrProvider::AzureAi, + "reducto" => OcrProvider::Reducto, + "vertex_ai" => OcrProvider::VertexAi, value => return Err(Error::InvalidProvider(value.to_string())), }; - match typed_provider { - OcrProvider::Mistral => Ok((provider.model.to_string(), OcrAdapterKind::Mistral)), + let adapter = match typed_provider { + OcrProvider::Mistral => OcrAdapterKind::Mistral, + OcrProvider::AzureAi if is_document_intelligence_model(provider.model) => { + OcrAdapterKind::AzureDocumentIntelligence + } + OcrProvider::AzureAi => OcrAdapterKind::AzureMistral, + OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-legacy") => { + OcrAdapterKind::ReductoLegacy + } + OcrProvider::Reducto if provider.model.eq_ignore_ascii_case("parse-v3") => { + OcrAdapterKind::ReductoV3 + } + OcrProvider::Reducto => { + return Err(Error::InvalidRequest(format!( + "unsupported Reducto OCR model: {}", + provider.model + ))); + } + OcrProvider::VertexAi if provider.model.to_ascii_lowercase().contains("deepseek") => { + OcrAdapterKind::VertexDeepSeek + } + OcrProvider::VertexAi => OcrAdapterKind::VertexMistral, + }; + Ok((provider.model.to_string(), adapter)) +} + +fn is_document_intelligence_model(model: &str) -> bool { + let model = model.to_ascii_lowercase(); + model.contains("doc-intelligence") || model.contains("documentintelligence") +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn provider_models_are_preserved_without_a_local_allowlist() { + let cases = [ + ("mistral/future-ocr-model", OcrAdapterKind::Mistral), + ("azure_ai/future-ocr-model", OcrAdapterKind::AzureMistral), + ]; + + for (qualified_model, expected_adapter) in cases { + let expected_model = qualified_model.split_once('/').unwrap().1; + let (model, adapter) = resolve_wire_adapter(qualified_model, None).unwrap(); + assert_eq!(model, expected_model); + assert_eq!(adapter, expected_adapter); + } + } + + #[test] + fn unknown_reducto_models_are_rejected() { + assert!(matches!( + resolve_wire_adapter("reducto/future-parse-model", None), + Err(Error::InvalidRequest(_)) + )); + } + + #[test] + fn known_protocol_models_still_select_specialized_adapters() { + let (model, adapter) = resolve_wire_adapter("reducto/parse-legacy", None).unwrap(); + assert_eq!(model, "parse-legacy"); + assert_eq!(adapter, OcrAdapterKind::ReductoLegacy); + + let (model, adapter) = + resolve_wire_adapter("azure_ai/doc-intelligence/prebuilt-layout", None).unwrap(); + assert_eq!(model, "doc-intelligence/prebuilt-layout"); + assert_eq!(adapter, OcrAdapterKind::AzureDocumentIntelligence); } } diff --git a/litellm-rust/crates/core/src/ocr/transformation.rs b/litellm-rust/crates/core/src/ocr/transformation.rs deleted file mode 100644 index ac4f10bf15b..00000000000 --- a/litellm-rust/crates/core/src/ocr/transformation.rs +++ /dev/null @@ -1,107 +0,0 @@ -use crate::Error; -use serde_json::{Map, Value}; - -use super::types::{LiteLLMOcrResponse, OcrRequestData}; - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum OcrAuthStrategy { - Bearer, - Header(&'static str), -} - -impl OcrAuthStrategy { - pub fn header_name(self) -> &'static str { - match self { - Self::Bearer => "authorization", - Self::Header(header_name) => header_name, - } - } -} - -#[derive(Clone, Copy, Debug, PartialEq, Eq)] -pub enum OcrResponseHandling { - Json, - AzureDocumentIntelligencePoll, -} - -pub trait OcrProviderConfig: Sync { - fn supported_ocr_params(&self) -> &'static [&'static str]; - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn map_ocr_params(&self, non_default_params: &Map) -> Map { - let mut mapped_params = Map::new(); - for (param, value) in non_default_params { - if self.supported_ocr_params().contains(¶m.as_str()) { - mapped_params.insert(param.clone(), value.clone()); - } - } - mapped_params - } - - fn transform_ocr_request( - &self, - model: &str, - document: Value, - optional_params: Map, - ) -> Result; - - fn transform_ocr_response( - &self, - model: &str, - response_json: Value, - ) -> Result; - - fn transform_ocr_response_with_params( - &self, - model: &str, - response_json: Value, - _optional_params: &Map, - ) -> Result { - self.transform_ocr_response(model, response_json) - } - - fn complete_url( - &self, - api_base: Option<&str>, - model: &str, - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; - - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result; - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn validate_environment( - &self, - headers: Vec<(String, String)>, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result, Error> { - let strategy = self.auth_strategy(); - if crate::http_utils::has_header(&headers, strategy.header_name()) { - return Ok(headers); - } - let api_key = self.resolve_api_key(api_key, env_lookup)?; - let auth_header = match strategy { - OcrAuthStrategy::Bearer => ("Authorization".to_string(), format!("Bearer {api_key}")), - OcrAuthStrategy::Header(name) => (name.to_string(), api_key), - }; - Ok(std::iter::once(auth_header).chain(headers).collect()) - } - - fn auth_strategy(&self) -> OcrAuthStrategy { - OcrAuthStrategy::Bearer - } - - fn requires_data_uri_document(&self) -> bool { - false - } - - fn response_handling(&self) -> OcrResponseHandling { - OcrResponseHandling::Json - } -} diff --git a/litellm-rust/crates/core/src/ocr/types.rs b/litellm-rust/crates/core/src/ocr/types.rs index 02b5330188c..06519f86c91 100644 --- a/litellm-rust/crates/core/src/ocr/types.rs +++ b/litellm-rust/crates/core/src/ocr/types.rs @@ -1,3 +1,4 @@ +use std::collections::BTreeMap; use std::sync::Arc; use std::time::Duration; @@ -7,14 +8,9 @@ use serde_json::{Map, Value}; use super::hooks::{NoopOcrHooks, OcrHooks}; use super::registry::{OcrAdapterKind, resolve_wire_adapter}; use crate::Error; +use crate::auth::InputSource; use crate::constants::OCR_HTTP_TIMEOUT_SECS; -#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] -pub struct OcrRequestData { - pub data: Value, - pub files: Option, -} - #[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] #[serde(tag = "type")] pub enum OcrDocument { @@ -32,6 +28,28 @@ pub enum OcrDocument { }, } +impl OcrDocument { + pub(crate) fn source(&self) -> &str { + match self { + Self::DocumentUrl { document_url, .. } => document_url, + Self::ImageUrl { image_url, .. } => image_url, + } + } + + pub(crate) fn with_source(self, source: String) -> Self { + match self { + Self::DocumentUrl { extra_fields, .. } => Self::DocumentUrl { + document_url: source, + extra_fields, + }, + Self::ImageUrl { extra_fields, .. } => Self::ImageUrl { + image_url: source, + extra_fields, + }, + } + } +} + #[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize, Deserialize)] #[serde(rename_all = "lowercase")] pub enum OcrResponseFormat { @@ -43,18 +61,28 @@ pub enum OcrResponseFormat { #[derive(Clone)] pub struct OcrConnection { pub api_key: Option, + pub api_key_source: InputSource, pub api_base: Option, + pub api_base_source: InputSource, pub extra_headers: Vec<(String, String)>, + pub extra_headers_source: InputSource, pub timeout: Duration, + pub max_download_bytes: u64, + pub poll_timeout: Duration, } impl Default for OcrConnection { fn default() -> Self { Self { api_key: None, + api_key_source: InputSource::Deployment, api_base: None, + api_base_source: InputSource::Deployment, extra_headers: Vec::new(), + extra_headers_source: InputSource::Deployment, timeout: Duration::from_secs(OCR_HTTP_TIMEOUT_SECS), + max_download_bytes: crate::constants::OCR_DOWNLOAD_MAX_BYTES, + poll_timeout: Duration::from_secs(crate::constants::OCR_POLL_TIMEOUT_SECS), } } } @@ -66,6 +94,7 @@ pub struct LiteLLMOcrRequest { pub hooks: Arc, pub litellm_call_id: Option, pub optional_params: Map, + pub input_sources: BTreeMap, pub(crate) adapter: OcrAdapterKind, } @@ -85,6 +114,7 @@ impl LiteLLMOcrRequest { hooks: Arc::new(NoopOcrHooks), litellm_call_id: None, optional_params, + input_sources: BTreeMap::new(), adapter: adapter_kind, }) } diff --git a/litellm-rust/crates/core/src/ocr/wire.rs b/litellm-rust/crates/core/src/ocr/wire.rs index f37c06a01ac..34d0a7d7b86 100644 --- a/litellm-rust/crates/core/src/ocr/wire.rs +++ b/litellm-rust/crates/core/src/ocr/wire.rs @@ -1,10 +1,12 @@ use crate::ocr::error::OcrRequestError; use crate::ocr::error::OcrResponseError; +use std::collections::BTreeMap; use std::time::Duration; use super::hooks::{OcrDuringCallRequest, OcrPreCallRequest}; use super::types::{LiteLLMOcrRequest, OcrConnection, OcrDocument}; use crate::Error; +use crate::auth::InputSource; use serde::{ Deserialize, de::{DeserializeOwned, IntoDeserializer}, @@ -28,6 +30,8 @@ pub struct OcrWireRequest { pub extra_headers: Option>, #[serde(default)] pub optional_params: Map, + #[serde(default)] + pub input_sources: BTreeMap, pub timeout_seconds: Option, } @@ -36,6 +40,9 @@ pub fn is_supported_request(model: &str, custom_llm_provider: Option<&str>) -> b } pub fn decode_request(wire: OcrWireRequest) -> Result { + let api_key_source = source_for(&wire.input_sources, "api_key"); + let api_base_source = source_for(&wire.input_sources, "api_base"); + let extra_headers_source = source_for(&wire.input_sources, "extra_headers"); let document = decode_request_value(wire.document, "document")?; let headers = wire .extra_headers @@ -67,16 +74,26 @@ pub fn decode_request(wire: OcrWireRequest) -> Result )?; let connection = OcrConnection { api_key: nonblank(wire.api_key), + api_key_source, api_base: nonblank(wire.api_base), + api_base_source, extra_headers: headers, + extra_headers_source, timeout: timeout.unwrap_or(defaults.timeout), + max_download_bytes: defaults.max_download_bytes, + poll_timeout: defaults.poll_timeout, }; Ok(LiteLLMOcrRequest { connection, + input_sources: wire.input_sources, ..request }) } +fn source_for(sources: &BTreeMap, name: &str) -> InputSource { + sources.get(name).copied().unwrap_or_default() +} + fn nonblank(value: Option) -> Option { value .map(|s| s.trim().to_string()) diff --git a/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs b/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs index f31b961e78a..3ed00b7cc5f 100644 --- a/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/anthropic/messages/transformation.rs @@ -1,3 +1,4 @@ +use crate::auth::error::MissingCredential; use crate::error::Error; use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy}; @@ -21,13 +22,7 @@ pub fn resolve_anthropic_api_key( non_empty(api_key) .map(str::to_string) .or_else(|| env_lookup(ANTHROPIC_API_KEY_ENV).filter(|value| !value.trim().is_empty())) - .ok_or_else(|| { - Error::Auth( - "Missing Anthropic API Key - Set `api_key` or the ANTHROPIC_API_KEY \ - environment variable" - .to_string(), - ) - }) + .ok_or_else(|| Error::from(crate::AuthError::from(MissingCredential::AnthropicApiKey))) } pub fn complete_anthropic_url( diff --git a/litellm-rust/crates/core/src/providers/azure_ai/auth/credential_provider_cache.rs b/litellm-rust/crates/core/src/providers/azure_ai/auth/credential_provider_cache.rs new file mode 100644 index 00000000000..297e4cc6502 --- /dev/null +++ b/litellm-rust/crates/core/src/providers/azure_ai/auth/credential_provider_cache.rs @@ -0,0 +1,43 @@ +use std::future::Future; +use std::sync::Arc; + +use azure_core::credentials::TokenCredential; +use moka::future::Cache; + +use crate::AuthError; + +#[derive(Clone, Debug, PartialEq, Eq, Hash)] +pub(crate) struct AzureCredentialProviderCacheKey { + pub(crate) mechanism: &'static str, + pub(crate) authority: String, + pub(crate) tenant_id: String, + pub(crate) client_id: String, + pub(crate) scope: String, + pub(crate) secret_identity: String, +} + +pub(crate) struct AzureCredentialProviderCache { + entries: Cache>, +} + +impl AzureCredentialProviderCache { + pub(crate) fn new(capacity: u64) -> Self { + Self { + entries: Cache::builder().max_capacity(capacity).build(), + } + } + + pub(crate) async fn get_or_create( + &self, + key: AzureCredentialProviderCacheKey, + create: F, + ) -> Result, AuthError> + where + F: Future, AuthError>>, + { + self.entries + .try_get_with(key, create) + .await + .map_err(|error| (*error).clone()) + } +} diff --git a/litellm-rust/crates/core/src/providers/azure_ai/auth/mod.rs b/litellm-rust/crates/core/src/providers/azure_ai/auth/mod.rs new file mode 100644 index 00000000000..33d007c1945 --- /dev/null +++ b/litellm-rust/crates/core/src/providers/azure_ai/auth/mod.rs @@ -0,0 +1,7 @@ +mod credential_provider_cache; +mod native; +mod resolve; +mod types; + +pub(crate) use resolve::AzureAuthService; +pub(crate) use types::AzureAuthInputs; diff --git a/litellm-rust/crates/core/src/providers/azure_ai/auth/native.rs b/litellm-rust/crates/core/src/providers/azure_ai/auth/native.rs new file mode 100644 index 00000000000..b8f19818d16 --- /dev/null +++ b/litellm-rust/crates/core/src/providers/azure_ai/auth/native.rs @@ -0,0 +1,702 @@ +use crate::auth::error::AuthConfigurationError; +use std::sync::Arc; +use std::time::{Duration, UNIX_EPOCH}; + +use azure_core::cloud::{CloudConfiguration, CustomConfiguration}; +use azure_core::credentials::{Secret, TokenCredential}; +use azure_core::http::ClientOptions; +use azure_identity::{ + ClientAssertion, ClientAssertionCredential, ClientAssertionCredentialOptions, + ClientSecretCredential, ClientSecretCredentialOptions, DeveloperToolsCredential, + ManagedIdentityCredential, ManagedIdentityCredentialOptions, UserAssignedId, + WorkloadIdentityCredential, WorkloadIdentityCredentialOptions, +}; +use sha2::{Digest, Sha256}; + +use crate::AuthError; +use crate::auth::{InputSource, ResolvedCredential, SecretValue, Sourced}; + +use super::credential_provider_cache::{ + AzureCredentialProviderCache, AzureCredentialProviderCacheKey, +}; + +#[derive(Clone, Debug)] +pub(crate) enum NativeAzureRequest { + ClientSecret { + tenant_id: Sourced, + client_id: Sourced, + client_secret: Sourced, + scope: Sourced, + authority: Option>, + }, + ClientAssertion { + tenant_id: Sourced, + client_id: Sourced, + assertion: Sourced, + assertion_identity: String, + scope: Sourced, + authority: Option>, + }, + WorkloadIdentity { + tenant_id: Sourced, + client_id: Sourced, + token_file_path: Sourced, + scope: Sourced, + authority: Option>, + }, + ManagedIdentity { + client_id: Option>, + scope: Sourced, + selection_source: InputSource, + }, + DeveloperTools { + scope: Sourced, + selection_source: InputSource, + }, +} + +#[derive(Clone, Debug)] +pub(crate) struct ValidatedAzureRequest { + request: NativeAzureRequest, + credential_source: InputSource, +} + +impl ValidatedAzureRequest { + pub(crate) fn new(request: NativeAzureRequest) -> Result { + validate_authority(&request)?; + let credential_source = validate_sources(&request)?; + Ok(Self { + request, + credential_source, + }) + } + + pub(crate) fn credential_source(&self) -> InputSource { + self.credential_source + } + + #[cfg(test)] + pub(super) fn kind(&self) -> &'static str { + match self.request { + NativeAzureRequest::ClientSecret { .. } => "client-secret", + NativeAzureRequest::ClientAssertion { .. } => "client-assertion", + NativeAzureRequest::WorkloadIdentity { .. } => "workload-identity", + NativeAzureRequest::ManagedIdentity { .. } => "managed-identity", + NativeAzureRequest::DeveloperTools { .. } => "developer-tools", + } + } +} + +pub(crate) struct NativeAzureTokenAcquirer { + cache: AzureCredentialProviderCache, + transport: Option, +} + +impl Default for NativeAzureTokenAcquirer { + fn default() -> Self { + Self::new(64) + } +} + +impl NativeAzureTokenAcquirer { + pub(crate) fn new(cache_capacity: u64) -> Self { + Self { + cache: AzureCredentialProviderCache::new(cache_capacity), + transport: None, + } + } + + #[cfg(test)] + pub(super) fn with_transport( + cache_capacity: u64, + transport: azure_core::http::Transport, + ) -> Self { + Self { + cache: AzureCredentialProviderCache::new(cache_capacity), + transport: Some(transport), + } + } + + pub(crate) async fn acquire( + &self, + request: ValidatedAzureRequest, + ) -> Result { + let scope = request.request.scope().to_string(); + let key = request.request.cache_key(); + let transport = self.transport.clone(); + let credential = self + .cache + .get_or_create( + key, + async move { build_credential(request.request, transport) }, + ) + .await?; + let token = credential + .get_token(&[scope.as_str()], None) + .await + .map_err(|error| AuthError::AzureTokenAcquisition(error.to_string()))?; + let expires_on = u64::try_from(token.expires_on.unix_timestamp()) + .ok() + .map(|seconds| UNIX_EPOCH + Duration::from_secs(seconds)); + + Ok(ResolvedCredential::AccessToken { + token: SecretValue::new(token.token.secret()), + expires_on, + }) + } +} + +impl NativeAzureRequest { + fn scope(&self) -> &str { + match self { + Self::ClientSecret { scope, .. } + | Self::ClientAssertion { scope, .. } + | Self::WorkloadIdentity { scope, .. } + | Self::ManagedIdentity { scope, .. } + | Self::DeveloperTools { scope, .. } => scope.value(), + } + } + + fn cache_key(&self) -> AzureCredentialProviderCacheKey { + match self { + Self::ClientSecret { + tenant_id, + client_id, + client_secret, + scope, + authority, + } => AzureCredentialProviderCacheKey { + mechanism: "client-secret", + authority: authority + .as_ref() + .map(|value| value.value().clone()) + .unwrap_or_default(), + tenant_id: tenant_id.value().clone(), + client_id: client_id.value().clone(), + scope: scope.value().clone(), + secret_identity: secret_digest(client_secret.value().expose()), + }, + Self::ClientAssertion { + tenant_id, + client_id, + assertion, + assertion_identity, + scope, + authority, + } => AzureCredentialProviderCacheKey { + mechanism: "client-assertion", + authority: authority + .as_ref() + .map(|value| value.value().clone()) + .unwrap_or_default(), + tenant_id: tenant_id.value().clone(), + client_id: client_id.value().clone(), + scope: scope.value().clone(), + secret_identity: format!( + "{assertion_identity}:{}", + secret_digest(assertion.value().expose()) + ), + }, + Self::WorkloadIdentity { + tenant_id, + client_id, + token_file_path, + scope, + authority, + } => AzureCredentialProviderCacheKey { + mechanism: "workload-identity", + authority: authority + .as_ref() + .map(|value| value.value().clone()) + .unwrap_or_default(), + tenant_id: tenant_id.value().clone(), + client_id: client_id.value().clone(), + scope: scope.value().clone(), + secret_identity: token_file_path.value().clone(), + }, + Self::ManagedIdentity { + client_id, scope, .. + } => AzureCredentialProviderCacheKey { + mechanism: "managed-identity", + authority: String::new(), + tenant_id: String::new(), + client_id: client_id + .as_ref() + .map(|value| value.value().clone()) + .unwrap_or_default(), + scope: scope.value().clone(), + secret_identity: String::new(), + }, + Self::DeveloperTools { scope, .. } => AzureCredentialProviderCacheKey { + mechanism: "developer-tools", + authority: String::new(), + tenant_id: String::new(), + client_id: String::new(), + scope: scope.value().clone(), + secret_identity: String::new(), + }, + } + } +} + +fn validate_authority(request: &NativeAzureRequest) -> Result<(), AuthError> { + let authority = match request { + NativeAzureRequest::ClientSecret { authority, .. } + | NativeAzureRequest::ClientAssertion { authority, .. } + | NativeAzureRequest::WorkloadIdentity { authority, .. } => authority.as_ref(), + NativeAzureRequest::ManagedIdentity { .. } | NativeAzureRequest::DeveloperTools { .. } => { + None + } + }; + let Some(authority) = authority else { + return Ok(()); + }; + let url = url::Url::parse(authority.value()) + .map_err(|_| AuthError::Configuration(AuthConfigurationError::InvalidAzureAuthority))?; + if url.scheme() != "https" + || url.host_str().is_none() + || !url.username().is_empty() + || url.password().is_some() + || url.query().is_some() + || url.fragment().is_some() + || !matches!(url.path(), "" | "/") + { + return Err(AuthError::Configuration( + AuthConfigurationError::InvalidAzureAuthority, + )); + } + Ok(()) +} + +fn validate_sources(request: &NativeAzureRequest) -> Result { + match request { + NativeAzureRequest::ClientSecret { + tenant_id, + client_id, + client_secret, + scope, + authority, + } => { + let identity_sources = [ + tenant_id.source(), + client_id.source(), + client_secret.source(), + ]; + let request_identity = identity_sources.contains(&InputSource::Request); + if request_identity + && !identity_sources + .iter() + .all(|source| *source == InputSource::Request) + { + return mixed_sources(); + } + if !request_identity && is_request_controlled(scope, authority.as_ref()) { + return mixed_sources(); + } + Ok(if request_identity { + InputSource::Request + } else { + trusted_source(&identity_sources) + }) + } + NativeAzureRequest::ClientAssertion { + tenant_id, + client_id, + assertion, + scope, + authority, + .. + } => trusted_only(&[ + tenant_id.source(), + client_id.source(), + assertion.source(), + scope.source(), + authority + .as_ref() + .map(Sourced::source) + .unwrap_or(InputSource::Environment), + ]), + NativeAzureRequest::WorkloadIdentity { + tenant_id, + client_id, + token_file_path, + scope, + authority, + } => trusted_only(&[ + tenant_id.source(), + client_id.source(), + token_file_path.source(), + scope.source(), + authority + .as_ref() + .map(Sourced::source) + .unwrap_or(InputSource::Environment), + ]), + NativeAzureRequest::ManagedIdentity { + client_id, + scope, + selection_source, + } => trusted_only(&[ + client_id + .as_ref() + .map(Sourced::source) + .unwrap_or(InputSource::Environment), + scope.source(), + *selection_source, + ]), + NativeAzureRequest::DeveloperTools { + scope, + selection_source, + } => trusted_only(&[scope.source(), *selection_source]), + } +} + +fn is_request_controlled(value: &Sourced, optional: Option<&Sourced>) -> bool { + value.source() == InputSource::Request + || optional.is_some_and(|value| value.source() == InputSource::Request) +} + +fn trusted_only(sources: &[InputSource]) -> Result { + if sources.contains(&InputSource::Request) { + return mixed_sources(); + } + Ok(trusted_source(sources)) +} + +fn trusted_source(sources: &[InputSource]) -> InputSource { + if sources.contains(&InputSource::Deployment) { + InputSource::Deployment + } else { + InputSource::Environment + } +} + +fn mixed_sources() -> Result { + Err(AuthError::Configuration( + AuthConfigurationError::MixedAzureCredentialSources, + )) +} + +fn build_credential( + request: NativeAzureRequest, + transport: Option, +) -> Result, AuthError> { + match request { + NativeAzureRequest::ClientSecret { + tenant_id, + client_id, + client_secret, + authority, + .. + } => ClientSecretCredential::new( + tenant_id.value(), + client_id.into_value(), + Secret::new(client_secret.value().expose().to_string()), + Some(ClientSecretCredentialOptions { + client_options: client_options(authority.map(Sourced::into_value), transport), + }), + ) + .map(|credential| credential as Arc), + NativeAzureRequest::ClientAssertion { + tenant_id, + client_id, + assertion, + authority, + .. + } => ClientAssertionCredential::new( + tenant_id.into_value(), + client_id.into_value(), + StaticAssertion(assertion.into_value()), + Some(ClientAssertionCredentialOptions { + client_options: client_options(authority.map(Sourced::into_value), transport), + }), + ) + .map(|credential| credential as Arc), + NativeAzureRequest::WorkloadIdentity { + tenant_id, + client_id, + token_file_path, + authority, + .. + } => WorkloadIdentityCredential::new(Some(WorkloadIdentityCredentialOptions { + credential_options: azure_identity::ClientAssertionCredentialOptions { + client_options: client_options(authority.map(Sourced::into_value), transport), + }, + client_id: Some(client_id.into_value()), + tenant_id: Some(tenant_id.into_value()), + token_file_path: Some(token_file_path.into_value().into()), + })) + .map(|credential| credential as Arc), + NativeAzureRequest::ManagedIdentity { client_id, .. } => { + ManagedIdentityCredential::new(Some(ManagedIdentityCredentialOptions { + user_assigned_id: client_id + .map(Sourced::into_value) + .map(UserAssignedId::ClientId), + client_options: client_options(None, transport), + })) + .map(|credential| credential as Arc) + } + NativeAzureRequest::DeveloperTools { .. } => DeveloperToolsCredential::new(None) + .map(|credential| credential as Arc), + } + .map_err(|error| { + AuthError::Configuration(AuthConfigurationError::AzureCredentialInitialization( + error.to_string(), + )) + }) +} + +fn client_options( + authority: Option, + transport: Option, +) -> ClientOptions { + let cloud = authority.map(|authority_host| { + let mut custom = CustomConfiguration::default(); + custom.authority_host = authority_host; + Arc::new(CloudConfiguration::from(custom)) + }); + ClientOptions { + cloud, + transport, + ..Default::default() + } +} + +fn secret_digest(secret: &str) -> String { + format!("{:x}", Sha256::digest(secret.as_bytes())) +} + +#[derive(Debug)] +struct StaticAssertion(SecretValue); + +impl ClientAssertion for StaticAssertion { + fn secret<'life0, 'life1, 'async_trait>( + &'life0 self, + _options: Option>, + ) -> std::pin::Pin< + Box> + Send + 'async_trait>, + > + where + 'life0: 'async_trait, + 'life1: 'async_trait, + Self: 'async_trait, + { + Box::pin(async move { Ok(self.0.expose().to_string()) }) + } +} + +#[cfg(test)] +mod tests { + use std::sync::{Arc, Mutex}; + + use azure_core::http::headers::Headers; + use azure_core::http::{AsyncRawResponse, HttpClient, Request, StatusCode, Transport}; + use azure_core::{Bytes, Result}; + + use super::{NativeAzureRequest, NativeAzureTokenAcquirer, ValidatedAzureRequest}; + use crate::auth::{InputSource, SecretValue, Sourced}; + + fn deployment(value: T) -> Sourced { + Sourced::new(value, InputSource::Deployment) + } + + fn sourced_client_secret( + credential_source: InputSource, + authority_source: InputSource, + authority: &str, + ) -> NativeAzureRequest { + NativeAzureRequest::ClientSecret { + tenant_id: Sourced::new("tenant".to_string(), credential_source), + client_id: Sourced::new("client".to_string(), credential_source), + client_secret: Sourced::new(SecretValue::new("secret"), credential_source), + scope: Sourced::new("scope".to_string(), InputSource::Environment), + authority: Some(Sourced::new(authority.to_string(), authority_source)), + } + } + + fn client_secret_request( + tenant: &str, + client: &str, + secret: &str, + scope: &str, + authority: &str, + ) -> ValidatedAzureRequest { + ValidatedAzureRequest::new(NativeAzureRequest::ClientSecret { + tenant_id: deployment(tenant.to_string()), + client_id: deployment(client.to_string()), + client_secret: deployment(SecretValue::new(secret)), + scope: deployment(scope.to_string()), + authority: Some(deployment(authority.to_string())), + }) + .unwrap() + } + + #[derive(Debug, Default)] + struct RecordingTokenClient { + requests: Mutex>, + } + + impl HttpClient for RecordingTokenClient { + fn execute_request<'life0, 'life1, 'async_trait>( + &'life0 self, + request: &'life1 Request, + ) -> std::pin::Pin< + Box> + Send + 'async_trait>, + > + where + 'life0: 'async_trait, + 'life1: 'async_trait, + Self: 'async_trait, + { + Box::pin(async move { + let body = Bytes::from(request.body()); + self.requests.lock().unwrap().push(( + request.url().to_string(), + String::from_utf8(body.to_vec()).unwrap(), + )); + Ok(AsyncRawResponse::from_bytes( + StatusCode::Ok, + Headers::new(), + r#"{"token_type":"Bearer","expires_in":3600,"ext_expires_in":3600,"access_token":"native-token"}"#, + )) + }) + } + } + + #[tokio::test] + async fn client_secret_uses_sdk_protocol_and_reuses_cached_credential() { + let transport = Arc::new(RecordingTokenClient::default()); + let acquirer = + NativeAzureTokenAcquirer::with_transport(4, Transport::new(transport.clone())); + let request = client_secret_request( + "tenant", + "client", + "secret", + "https://service.test/.default", + "https://login.test", + ); + + let first = acquirer.acquire(request.clone()).await.unwrap(); + let second = acquirer.acquire(request).await.unwrap(); + + assert_eq!(first.secret().expose(), "native-token"); + assert_eq!(second.secret().expose(), "native-token"); + let requests = transport.requests.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert_eq!(requests[0].0, "https://login.test/tenant/oauth2/v2.0/token"); + assert!(requests[0].1.contains("client_id=client")); + assert!(requests[0].1.contains("client_secret=secret")); + assert!( + requests[0] + .1 + .contains("scope=https%3A%2F%2Fservice.test%2F.default") + ); + } + + #[tokio::test] + async fn credential_provider_cache_isolates_every_client_secret_identity_field() { + let transport = Arc::new(RecordingTokenClient::default()); + let acquirer = + NativeAzureTokenAcquirer::with_transport(16, Transport::new(transport.clone())); + let request = client_secret_request; + let base = request("tenant", "client", "secret", "scope", "https://login.test"); + let variants = [ + base.clone(), + request( + "other-tenant", + "client", + "secret", + "scope", + "https://login.test", + ), + request( + "tenant", + "other-client", + "secret", + "scope", + "https://login.test", + ), + request( + "tenant", + "client", + "other-secret", + "scope", + "https://login.test", + ), + request( + "tenant", + "client", + "secret", + "other-scope", + "https://login.test", + ), + request( + "tenant", + "client", + "secret", + "scope", + "https://other-login.test", + ), + ]; + + acquirer.acquire(base.clone()).await.unwrap(); + acquirer.acquire(base).await.unwrap(); + for request in variants.into_iter().skip(1) { + acquirer.acquire(request).await.unwrap(); + } + + assert_eq!(transport.requests.lock().unwrap().len(), 6); + } + + #[test] + fn request_authority_requires_request_owned_client_secret_identity() { + let error = ValidatedAzureRequest::new(sourced_client_secret( + InputSource::Deployment, + InputSource::Request, + "https://login.example", + )) + .unwrap_err(); + + assert!(matches!( + error, + crate::AuthError::Configuration( + crate::auth::error::AuthConfigurationError::MixedAzureCredentialSources + ) + )); + } + + #[test] + fn request_owned_client_secret_identity_can_select_custom_authority() { + let request = ValidatedAzureRequest::new(sourced_client_secret( + InputSource::Request, + InputSource::Request, + "https://login.example", + )) + .unwrap(); + + assert_eq!(request.credential_source(), InputSource::Request); + } + + #[test] + fn authority_is_restricted_to_an_https_origin() { + for authority in [ + "http://login.example", + "https://user@login.example", + "https://login.example/tenant", + "https://login.example?target=other", + ] { + let error = ValidatedAzureRequest::new(sourced_client_secret( + InputSource::Deployment, + InputSource::Deployment, + authority, + )) + .unwrap_err(); + assert!(matches!( + error, + crate::AuthError::Configuration( + crate::auth::error::AuthConfigurationError::InvalidAzureAuthority + ) + )); + } + } +} diff --git a/litellm-rust/crates/core/src/providers/azure_ai/auth/resolve.rs b/litellm-rust/crates/core/src/providers/azure_ai/auth/resolve.rs new file mode 100644 index 00000000000..025dd4f8740 --- /dev/null +++ b/litellm-rust/crates/core/src/providers/azure_ai/auth/resolve.rs @@ -0,0 +1,683 @@ +use crate::AuthError; +use crate::auth::error::AuthConfigurationError; +use crate::auth::{ + CredentialFileRef, CredentialLookup, CredentialRef, InputSource, ResolvedCredential, + SecretValue, Sourced, TokenProviderHandle, +}; +use std::future::Future; +use std::pin::Pin; +use std::sync::Arc; + +use super::native::{NativeAzureRequest, NativeAzureTokenAcquirer, ValidatedAzureRequest}; +use super::types::{AzureAuthInputs, AzureCredentialType, ConfigValue, DEFAULT_AZURE_SCOPE}; + +const AZURE_AD_TOKEN_ENV: &str = "AZURE_AD_TOKEN"; +const AZURE_TENANT_ID_ENV: &str = "AZURE_TENANT_ID"; +const AZURE_CLIENT_ID_ENV: &str = "AZURE_CLIENT_ID"; +const AZURE_CLIENT_SECRET_ENV: &str = "AZURE_CLIENT_SECRET"; +const AZURE_SCOPE_ENV: &str = "AZURE_SCOPE"; +const AZURE_AUTHORITY_HOST_ENV: &str = "AZURE_AUTHORITY_HOST"; +const AZURE_CREDENTIAL_ENV: &str = "AZURE_CREDENTIAL"; +const AZURE_FEDERATED_TOKEN_FILE_ENV: &str = "AZURE_FEDERATED_TOKEN_FILE"; + +#[derive(Clone, Debug)] +pub(crate) enum AzureCredentialPlan { + Supplied(Sourced), + Caller(TokenProviderHandle), + Oidc { + reference: Sourced, + tenant_id: Sourced, + client_id: Sourced, + scope: Sourced, + authority: Option>, + }, + Native(ValidatedAzureRequest), + Chain(Vec), + Missing, +} + +/// Rust counterpart to Python's `get_azure_ad_token`, not `BaseAzureLLM`. +pub(crate) struct AzureAuthService { + native: Arc, +} + +trait AzureTokenAcquirer: Send + Sync { + fn acquire( + &self, + request: ValidatedAzureRequest, + ) -> Pin> + Send + '_>>; +} + +impl AzureTokenAcquirer for NativeAzureTokenAcquirer { + fn acquire( + &self, + request: ValidatedAzureRequest, + ) -> Pin> + Send + '_>> { + Box::pin(NativeAzureTokenAcquirer::acquire(self, request)) + } +} + +impl Default for AzureAuthService { + fn default() -> Self { + Self { + native: Arc::new(NativeAzureTokenAcquirer::default()), + } + } +} + +impl AzureAuthService { + #[cfg(test)] + fn with_acquirer(native: Arc) -> Self { + Self { native } + } + + pub(crate) async fn get_azure_ad_token( + &self, + inputs: &AzureAuthInputs, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + ) -> Result>, AuthError> { + match select_auth_plan(inputs, env_lookup)? { + AzureCredentialPlan::Supplied(credential) => Ok(Some(credential)), + AzureCredentialPlan::Caller(caller) => { + let credential = caller.acquire().await?; + if credential.secret().expose().is_empty() { + return Err(AuthError::EmptyAzureToken); + } + Ok(Some(Sourced::new(credential, InputSource::Deployment))) + } + AzureCredentialPlan::Oidc { + reference, + tenant_id, + client_id, + scope, + authority, + } => { + let assertion = resolve_reference(inputs, env_lookup, reference.value()) + .await? + .ok_or(AuthError::UnresolvedOidcReference)?; + let request = ValidatedAzureRequest::new(NativeAzureRequest::ClientAssertion { + tenant_id, + client_id, + assertion: Sourced::new(assertion, reference.source()), + assertion_identity: format!("{:?}", reference.value()), + scope, + authority, + })?; + let source = request.credential_source(); + self.native + .acquire(request) + .await + .map(|credential| Sourced::new(credential, source)) + .map(Some) + } + AzureCredentialPlan::Native(request) => { + let source = request.credential_source(); + self.native + .acquire(request) + .await + .map(|credential| Some(Sourced::new(credential, source))) + } + AzureCredentialPlan::Chain(requests) => { + let mut failures = Vec::new(); + for request in requests { + let source = request.credential_source(); + match self.native.acquire(request).await { + Ok(credential) => return Ok(Some(Sourced::new(credential, source))), + Err(error) => failures.push(error), + } + } + Err(AuthError::CredentialChain(failures)) + } + AzureCredentialPlan::Missing => Ok(None), + } + } +} + +pub(crate) fn select_auth_plan( + inputs: &AzureAuthInputs, + env_lookup: &dyn Fn(&str) -> Option, +) -> Result { + let token = configured_secret(&inputs.azure_ad_token, AZURE_AD_TOKEN_ENV, env_lookup); + let tenant_id = configured_string(&inputs.tenant_id, AZURE_TENANT_ID_ENV, env_lookup); + let client_id = configured_string(&inputs.client_id, AZURE_CLIENT_ID_ENV, env_lookup); + let client_secret = + configured_secret(&inputs.client_secret, AZURE_CLIENT_SECRET_ENV, env_lookup); + let scope = configured_string(&inputs.azure_scope, AZURE_SCOPE_ENV, env_lookup) + .unwrap_or_else(|| Sourced::new(DEFAULT_AZURE_SCOPE.to_string(), InputSource::Environment)); + let authority = configured_string( + &inputs.azure_authority_host, + AZURE_AUTHORITY_HOST_ENV, + env_lookup, + ); + let selector = configured_string(&inputs.azure_credential, AZURE_CREDENTIAL_ENV, env_lookup) + .map(|value| { + value + .value() + .parse::() + .map(|selector| Sourced::new(selector, value.source())) + }) + .transpose() + .map_err(|_| AuthError::Configuration(AuthConfigurationError::InvalidAzureSelector))?; + let federated_token_file = configured_string( + &inputs.federated_token_file, + AZURE_FEDERATED_TOKEN_FILE_ENV, + env_lookup, + ); + + if inputs.azure_ad_token_provider.is_none() + && let (Some(tenant_id), Some(client_id), Some(client_secret)) = + (tenant_id.clone(), client_id.clone(), client_secret) + { + return Ok(AzureCredentialPlan::Native(ValidatedAzureRequest::new( + NativeAzureRequest::ClientSecret { + tenant_id, + client_id, + client_secret, + scope, + authority, + }, + )?)); + } + + if let (Some(reference), Some(tenant_id), Some(client_id)) = ( + oidc_reference(&token)?, + tenant_id.clone(), + client_id.clone(), + ) { + return Ok(AzureCredentialPlan::Oidc { + reference, + tenant_id, + client_id, + scope, + authority, + }); + } + + if let Some(caller) = &inputs.azure_ad_token_provider { + return Ok(AzureCredentialPlan::Caller(caller.clone())); + } + + if let Some(token) = token { + return Ok(AzureCredentialPlan::Supplied(token.map(|token| { + ResolvedCredential::AccessToken { + token, + expires_on: None, + } + }))); + } + + if !*inputs.enable_azure_ad_token_refresh.value() && selector.is_none() { + return Ok(AzureCredentialPlan::Missing); + } + + select_native_plan( + selector, + tenant_id, + client_id, + federated_token_file, + scope, + authority, + inputs.enable_azure_ad_token_refresh.source(), + ) +} + +fn select_native_plan( + selector: Option>, + tenant_id: Option>, + client_id: Option>, + federated_token_file: Option>, + scope: Sourced, + authority: Option>, + refresh_source: InputSource, +) -> Result { + let selected = selector.unwrap_or_else(|| { + Sourced::new( + { + if federated_token_file.is_some() { + AzureCredentialType::DefaultAzureCredential + } else if client_id.is_some() { + AzureCredentialType::ManagedIdentityCredential + } else { + AzureCredentialType::DefaultAzureCredential + } + }, + refresh_source, + ) + }); + let selection_source = selected.source(); + + match selected.into_value() { + AzureCredentialType::ClientSecretCredential => Err(AuthError::Configuration( + AuthConfigurationError::MissingClientSecretFields, + )), + AzureCredentialType::WorkloadIdentityCredential => { + Ok(AzureCredentialPlan::Native(ValidatedAzureRequest::new( + workload_request(tenant_id, client_id, federated_token_file, scope, authority)?, + )?)) + } + AzureCredentialType::ManagedIdentityCredential => Ok(AzureCredentialPlan::Native( + ValidatedAzureRequest::new(NativeAzureRequest::ManagedIdentity { + client_id, + scope, + selection_source, + })?, + )), + AzureCredentialType::DefaultAzureCredential => { + let workload = match (tenant_id, client_id.clone(), federated_token_file) { + (Some(tenant_id), Some(client_id), Some(token_file_path)) => { + Some(NativeAzureRequest::WorkloadIdentity { + tenant_id, + client_id, + token_file_path, + scope: scope.clone(), + authority, + }) + } + _ => None, + }; + Ok(AzureCredentialPlan::Chain( + workload + .into_iter() + .chain(std::iter::once(NativeAzureRequest::ManagedIdentity { + client_id, + scope: scope.clone(), + selection_source, + })) + .chain(std::iter::once(NativeAzureRequest::DeveloperTools { + scope, + selection_source, + })) + .map(ValidatedAzureRequest::new) + .collect::, _>>()?, + )) + } + AzureCredentialType::DeploymentIdentityCredential => { + let workload = match (tenant_id, client_id.clone(), federated_token_file) { + (Some(tenant_id), Some(client_id), Some(token_file_path)) => { + Some(NativeAzureRequest::WorkloadIdentity { + tenant_id, + client_id, + token_file_path, + scope: scope.clone(), + authority, + }) + } + _ => None, + }; + let user_assigned = client_id.map(|client_id| NativeAzureRequest::ManagedIdentity { + client_id: Some(client_id), + scope: scope.clone(), + selection_source, + }); + Ok(AzureCredentialPlan::Chain( + workload + .into_iter() + .chain(user_assigned) + .chain(std::iter::once(NativeAzureRequest::ManagedIdentity { + client_id: None, + scope, + selection_source, + })) + .map(ValidatedAzureRequest::new) + .collect::, _>>()?, + )) + } + } +} + +fn workload_request( + tenant_id: Option>, + client_id: Option>, + token_file_path: Option>, + scope: Sourced, + authority: Option>, +) -> Result { + Ok(NativeAzureRequest::WorkloadIdentity { + tenant_id: tenant_id.ok_or(AuthError::Configuration( + AuthConfigurationError::MissingWorkloadTenant, + ))?, + client_id: client_id.ok_or(AuthError::Configuration( + AuthConfigurationError::MissingWorkloadClient, + ))?, + token_file_path: token_file_path.ok_or(AuthError::Configuration( + AuthConfigurationError::MissingWorkloadTokenFile, + ))?, + scope, + authority, + }) +} + +fn configured_string( + configured: &ConfigValue, + environment_name: &str, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option> { + configured + .as_value() + .filter(|value| !value.value().is_empty()) + .cloned() + .or_else(|| { + env_lookup(environment_name) + .filter(|value| !value.is_empty()) + .map(|value| Sourced::new(value, InputSource::Environment)) + }) +} + +fn configured_secret( + configured: &ConfigValue, + environment_name: &str, + env_lookup: &dyn Fn(&str) -> Option, +) -> Option> { + configured + .as_value() + .filter(|value| !value.value().expose().is_empty()) + .cloned() + .or_else(|| { + env_lookup(environment_name) + .filter(|value| !value.is_empty()) + .map(|value| Sourced::new(SecretValue::new(value), InputSource::Environment)) + }) +} + +async fn resolve_reference( + inputs: &AzureAuthInputs, + env_lookup: &(dyn Fn(&str) -> Option + Sync), + reference: &CredentialRef, +) -> Result, AuthError> { + let lookup = match reference { + CredentialRef::Explicit(secret) => return Ok(Some(secret.clone())), + CredentialRef::Env(name) => env_lookup(name) + .filter(|value| !value.is_empty()) + .map(SecretValue::new) + .map_or(CredentialLookup::Missing, CredentialLookup::Found), + CredentialRef::None => return Ok(None), + CredentialRef::File(_) | CredentialRef::Request(_) | CredentialRef::Host(_) => { + let resolver = inputs + .credential_resolver + .as_ref() + .ok_or(AuthError::Configuration( + AuthConfigurationError::MissingHostResolver, + ))?; + resolver.resolve(reference).await? + } + }; + Ok(match lookup { + CredentialLookup::Found(secret) => Some(secret), + CredentialLookup::Missing | CredentialLookup::Declined => None, + }) +} + +fn oidc_reference( + token: &Option>, +) -> Result>, AuthError> { + let Some(token) = token.as_ref() else { + return Ok(None); + }; + let value = token.value().expose(); + if token.source() == InputSource::Request && value.starts_with("oidc/") { + return Err(AuthError::Configuration( + AuthConfigurationError::RequestAzureCredentialReference, + )); + } + if let Some(name) = value.strip_prefix("oidc/env/") { + return non_empty_reference(name, "OIDC environment reference") + .map(CredentialRef::Env) + .map(|reference| Sourced::new(reference, token.source())) + .map(Some); + } + if let Some(name) = value.strip_prefix("oidc/env_path/") { + return non_empty_reference(name, "OIDC environment path reference") + .map(|name| CredentialRef::File(CredentialFileRef::EnvironmentVariable(name))) + .map(|reference| Sourced::new(reference, token.source())) + .map(Some); + } + if let Some(path) = value.strip_prefix("oidc/file/") { + let path = non_empty_reference(path, "OIDC file reference")?; + return Ok(Some(Sourced::new( + CredentialRef::File(CredentialFileRef::Path(path.into())), + token.source(), + ))); + } + if value.starts_with("oidc/") { + return Err(AuthError::Configuration( + AuthConfigurationError::UnsupportedOidcReference, + )); + } + Ok(None) +} + +fn non_empty_reference(value: &str, kind: &str) -> Result { + if value.is_empty() { + return Err(AuthError::Configuration( + AuthConfigurationError::EmptyReference(kind.to_string()), + )); + } + Ok(value.to_string()) +} + +#[cfg(test)] +mod tests { + use std::future::Future; + use std::sync::{Arc, Mutex}; + + use serde_json::json; + + use super::{ + AzureAuthService, AzureCredentialPlan, AzureTokenAcquirer, oidc_reference, + resolve_reference, select_auth_plan, + }; + use crate::AuthError; + use crate::auth::ResolvedCredential; + use crate::auth::{ + CredentialFileRef, CredentialLookup, CredentialLookupFuture, CredentialRef, + CredentialResolver, CredentialResolverHandle, InputSource, SecretValue, Sourced, + }; + use crate::providers::azure_ai::auth::native::ValidatedAzureRequest; + use crate::providers::azure_ai::auth::types::AzureAuthInputs; + + #[derive(Debug)] + struct FileResolver; + + struct ChainAcquirer { + requests: Mutex>, + succeed_on: Option<&'static str>, + } + + impl AzureTokenAcquirer for ChainAcquirer { + fn acquire( + &self, + request: ValidatedAzureRequest, + ) -> std::pin::Pin< + Box> + Send + '_>, + > { + let kind = request.kind(); + self.requests.lock().unwrap().push(kind); + Box::pin(async move { + if self.succeed_on == Some(kind) { + Ok(ResolvedCredential::AccessToken { + token: SecretValue::new("chain-token"), + expires_on: None, + }) + } else { + Err(AuthError::AzureTokenAcquisition(format!("{kind} failed"))) + } + }) + } + } + + impl CredentialResolver for FileResolver { + fn resolve<'a>(&'a self, reference: &'a CredentialRef) -> CredentialLookupFuture<'a> { + Box::pin(async move { + Ok(match reference { + CredentialRef::File(CredentialFileRef::Path(path)) + if path == std::path::Path::new("/run/secrets/assertion") => + { + CredentialLookup::Found(SecretValue::new("rotated-assertion")) + } + _ => CredentialLookup::Declined, + }) + }) + } + } + + #[test] + fn null_and_empty_values_fall_back_to_environment() { + let params = json!({"tenant_id": null, "client_id": "", "client_secret": null}); + let inputs = AzureAuthInputs::from_optional_params(params.as_object().unwrap()).unwrap(); + let plan = select_auth_plan(&inputs, &|name| match name { + "AZURE_TENANT_ID" => Some("tenant".to_string()), + "AZURE_CLIENT_ID" => Some("client".to_string()), + "AZURE_CLIENT_SECRET" => Some("secret".to_string()), + _ => None, + }) + .unwrap(); + + assert!(matches!(plan, AzureCredentialPlan::Native(_))); + } + + #[test] + fn supplied_token_does_not_require_refresh() { + let params = json!({"azure_ad_token": "token"}); + let inputs = AzureAuthInputs::from_optional_params(params.as_object().unwrap()).unwrap(); + + assert!(matches!( + select_auth_plan(&inputs, &|_| None).unwrap(), + AzureCredentialPlan::Supplied(_) + )); + } + + #[test] + fn oidc_reference_is_deferred() { + let params = json!({ + "azure_ad_token": "oidc/env/ASSERTION", + "tenant_id": "tenant", + "client_id": "client" + }); + let inputs = AzureAuthInputs::from_optional_params(params.as_object().unwrap()).unwrap(); + + assert!(matches!( + select_auth_plan(&inputs, &|_| None).unwrap(), + AzureCredentialPlan::Oidc { + reference, + .. + } if reference.value() == &CredentialRef::Env("ASSERTION".to_string()) + )); + } + + #[test] + fn oidc_file_location_is_typed_before_resolution() { + assert_eq!( + oidc_reference(&Some(Sourced::new( + SecretValue::new("oidc/file//run/secrets/assertion"), + InputSource::Deployment, + ))) + .unwrap() + .map(Sourced::into_value), + Some(CredentialRef::File(CredentialFileRef::Path( + "/run/secrets/assertion".into() + ))) + ); + } + + #[test] + fn unsupported_oidc_reference_is_rejected_during_plan_creation() { + let error = oidc_reference(&Some(Sourced::new( + SecretValue::new("oidc/vault/assertion"), + InputSource::Deployment, + ))) + .expect_err("unsupported backend must fail validation"); + + assert!(error.to_string().contains("unsupported OIDC reference")); + } + + #[test] + fn request_oidc_reference_is_rejected_before_lookup() { + let params = json!({ + "azure_ad_token": "oidc/env/ASSERTION", + "tenant_id": "tenant", + "client_id": "client" + }); + let sources = std::collections::BTreeMap::from([ + ("azure_ad_token".to_string(), InputSource::Request), + ("tenant_id".to_string(), InputSource::Request), + ("client_id".to_string(), InputSource::Request), + ]); + let inputs = + AzureAuthInputs::from_sourced_optional_params(params.as_object().unwrap(), &sources) + .unwrap(); + + let error = select_auth_plan(&inputs, &|name| { + assert_ne!(name, "ASSERTION"); + None + }) + .unwrap_err(); + + assert!(matches!( + error, + AuthError::Configuration( + crate::auth::error::AuthConfigurationError::RequestAzureCredentialReference + ) + )); + } + + #[tokio::test] + async fn host_resolver_owns_file_access() { + let inputs = AzureAuthInputs { + credential_resolver: Some(CredentialResolverHandle::new(Arc::new(FileResolver))), + ..AzureAuthInputs::default() + }; + let reference = + CredentialRef::File(CredentialFileRef::Path("/run/secrets/assertion".into())); + + let resolved = resolve_reference(&inputs, &|_| None, &reference) + .await + .unwrap(); + + assert_eq!(resolved, Some(SecretValue::new("rotated-assertion"))); + } + + #[tokio::test] + async fn default_chain_uses_declared_order_and_stops_after_success() { + let acquirer = Arc::new(ChainAcquirer { + requests: Mutex::new(Vec::new()), + succeed_on: Some("developer-tools"), + }); + let service = AzureAuthService::with_acquirer(acquirer.clone()); + let inputs = AzureAuthInputs { + enable_azure_ad_token_refresh: Sourced::new(true, InputSource::Deployment), + ..Default::default() + }; + + let credential = service + .get_azure_ad_token(&inputs, &|_| None) + .await + .unwrap() + .unwrap(); + + assert_eq!(credential.value().secret().expose(), "chain-token"); + assert_eq!( + *acquirer.requests.lock().unwrap(), + ["managed-identity", "developer-tools"] + ); + } + + #[tokio::test] + async fn chain_reports_each_acquisition_failure() { + let acquirer = Arc::new(ChainAcquirer { + requests: Mutex::new(Vec::new()), + succeed_on: None, + }); + let service = AzureAuthService::with_acquirer(acquirer); + let inputs = AzureAuthInputs { + enable_azure_ad_token_refresh: Sourced::new(true, InputSource::Deployment), + ..Default::default() + }; + + let error = service + .get_azure_ad_token(&inputs, &|_| None) + .await + .unwrap_err(); + + assert!(matches!(error, AuthError::CredentialChain(errors) if errors.len() == 2)); + } +} diff --git a/litellm-rust/crates/core/src/providers/azure_ai/auth/types.rs b/litellm-rust/crates/core/src/providers/azure_ai/auth/types.rs new file mode 100644 index 00000000000..f15d526d945 --- /dev/null +++ b/litellm-rust/crates/core/src/providers/azure_ai/auth/types.rs @@ -0,0 +1,195 @@ +use crate::auth::error::AuthConfigurationError; +use serde_json::{Map, Value}; +use std::collections::BTreeMap; +use strum::EnumString; + +use crate::AuthError; +use crate::auth::{ + CredentialResolverHandle, InputSource, SecretValue, Sourced, TokenProviderHandle, +}; + +pub const DEFAULT_AZURE_SCOPE: &str = "https://cognitiveservices.azure.com/.default"; + +#[derive(Clone, Debug, Default, PartialEq, Eq)] +pub enum ConfigValue { + #[default] + Absent, + ExplicitNone(InputSource), + Value(Sourced), +} + +impl ConfigValue { + pub fn as_value(&self) -> Option<&Sourced> { + match self { + Self::Value(value) => Some(value), + Self::Absent | Self::ExplicitNone(_) => None, + } + } +} + +#[derive(Clone, Copy, Debug, EnumString, PartialEq, Eq, Hash)] +#[allow(clippy::enum_variant_names)] +pub enum AzureCredentialType { + ClientSecretCredential, + ManagedIdentityCredential, + DefaultAzureCredential, + DeploymentIdentityCredential, + WorkloadIdentityCredential, +} + +#[derive(Clone, Debug, Default)] +pub struct AzureAuthInputs { + pub azure_ad_token: ConfigValue, + pub azure_ad_token_provider: Option, + pub credential_resolver: Option, + pub tenant_id: ConfigValue, + pub client_id: ConfigValue, + pub client_secret: ConfigValue, + pub azure_scope: ConfigValue, + pub azure_authority_host: ConfigValue, + pub azure_credential: ConfigValue, + pub federated_token_file: ConfigValue, + pub enable_azure_ad_token_refresh: Sourced, +} + +impl AzureAuthInputs { + #[cfg(test)] + pub fn from_optional_params(params: &Map) -> Result { + Self::from_sourced_optional_params(params, &BTreeMap::new()) + } + + pub fn from_sourced_optional_params( + params: &Map, + sources: &BTreeMap, + ) -> Result { + Ok(Self { + azure_ad_token: secret_config(params, sources, "azure_ad_token")?, + azure_ad_token_provider: None, + credential_resolver: None, + tenant_id: string_config(params, sources, "tenant_id")?, + client_id: string_config(params, sources, "client_id")?, + client_secret: secret_config(params, sources, "client_secret")?, + azure_scope: string_config(params, sources, "azure_scope")?, + azure_authority_host: string_config(params, sources, "azure_authority_host")?, + azure_credential: string_config(params, sources, "azure_credential")?, + federated_token_file: string_config(params, sources, "azure_federated_token_file")?, + enable_azure_ad_token_refresh: Sourced::new( + params + .get("enable_azure_ad_token_refresh") + .and_then(Value::as_bool) + .unwrap_or(false), + source_for(sources, "enable_azure_ad_token_refresh"), + ), + }) + } +} + +fn string_config( + params: &Map, + sources: &BTreeMap, + name: &str, +) -> Result, AuthError> { + let source = source_for(sources, name); + match params.get(name) { + None => Ok(ConfigValue::Absent), + Some(Value::Null) => Ok(ConfigValue::ExplicitNone(source)), + Some(Value::String(value)) => Ok(ConfigValue::Value(Sourced::new(value.clone(), source))), + Some(_) => Err(AuthError::Configuration( + AuthConfigurationError::InvalidFieldType(name.to_string()), + )), + } +} + +fn secret_config( + params: &Map, + sources: &BTreeMap, + name: &str, +) -> Result, AuthError> { + Ok(match string_config(params, sources, name)? { + ConfigValue::Absent => ConfigValue::Absent, + ConfigValue::ExplicitNone(source) => ConfigValue::ExplicitNone(source), + ConfigValue::Value(value) => ConfigValue::Value(value.map(SecretValue::new)), + }) +} + +fn source_for(sources: &BTreeMap, name: &str) -> InputSource { + sources.get(name).copied().unwrap_or_default() +} + +#[cfg(test)] +mod tests { + use serde_json::json; + + use std::collections::BTreeMap; + + use super::{AzureAuthInputs, AzureCredentialType, ConfigValue}; + use crate::auth::{InputSource, Sourced}; + + #[test] + fn selector_parsing_is_exact() { + assert_eq!( + "ClientSecretCredential".parse::(), + Ok(AzureCredentialType::ClientSecretCredential) + ); + assert!( + "clientsecretcredential" + .parse::() + .is_err() + ); + } + + #[test] + fn defaults_preserve_absence() { + let inputs = AzureAuthInputs::default(); + + assert_eq!(inputs.tenant_id, ConfigValue::Absent); + assert_eq!(inputs.azure_ad_token, ConfigValue::Absent); + } + + #[test] + fn parsing_distinguishes_null_empty_and_absent() { + let params = json!({"tenant_id": null, "client_id": ""}); + let inputs = AzureAuthInputs::from_optional_params(params.as_object().unwrap()).unwrap(); + + assert_eq!( + inputs.tenant_id, + ConfigValue::ExplicitNone(InputSource::Deployment) + ); + assert_eq!( + inputs.client_id, + ConfigValue::Value(Sourced::new(String::new(), InputSource::Deployment)) + ); + assert_eq!(inputs.client_secret, ConfigValue::Absent); + } + + #[test] + fn parsing_preserves_trusted_input_sources() { + let params = json!({"tenant_id": "tenant", "client_secret": null}); + let sources = BTreeMap::from([ + ("tenant_id".to_string(), InputSource::Request), + ("client_secret".to_string(), InputSource::Request), + ]); + let inputs = + AzureAuthInputs::from_sourced_optional_params(params.as_object().unwrap(), &sources) + .unwrap(); + + assert_eq!( + inputs.tenant_id, + ConfigValue::Value(Sourced::new("tenant".to_string(), InputSource::Request)) + ); + assert_eq!( + inputs.client_secret, + ConfigValue::ExplicitNone(InputSource::Request) + ); + } + + #[test] + fn debug_does_not_expose_secrets() { + let params = json!({"azure_ad_token": "token-value", "client_secret": "secret-value"}); + let inputs = AzureAuthInputs::from_optional_params(params.as_object().unwrap()).unwrap(); + let debug = format!("{inputs:?}"); + + assert!(!debug.contains("token-value")); + assert!(!debug.contains("secret-value")); + } +} diff --git a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs index b8ca10461fb..585b34f393f 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/messages/transformation.rs @@ -1,3 +1,4 @@ +use crate::auth::error::MissingCredential; use crate::error::Error; use crate::messages::transformation::{AnthropicMessagesProviderConfig, MessagesAuthStrategy}; use crate::messages::types::{ @@ -32,12 +33,7 @@ pub fn resolve_azure_api_key( non_empty(api_key) .map(str::to_string) .or_else(|| env_lookup(AZURE_API_KEY_ENV).filter(|value| !value.trim().is_empty())) - .ok_or_else(|| { - Error::Auth( - "Missing Azure API Key - Set `api_key` or the AZURE_API_KEY environment variable" - .to_string(), - ) - }) + .ok_or_else(|| Error::from(crate::AuthError::from(MissingCredential::AzureApiKey))) } pub fn complete_azure_anthropic_url( @@ -47,13 +43,7 @@ pub fn complete_azure_anthropic_url( let api_base = non_empty(api_base) .map(str::to_string) .or_else(|| env_lookup(AZURE_API_BASE_ENV).filter(|value| !value.trim().is_empty())) - .ok_or_else(|| { - Error::Auth( - "Missing Azure API Base - Set `api_base` or the AZURE_API_BASE environment variable. \ - Expected format: https://.services.ai.azure.com/anthropic" - .to_string(), - ) - })?; + .ok_or_else(|| Error::from(crate::AuthError::from(MissingCredential::AzureApiBase)))?; let api_base = api_base.trim_end_matches('/'); diff --git a/litellm-rust/crates/core/src/providers/azure_ai/mod.rs b/litellm-rust/crates/core/src/providers/azure_ai/mod.rs index 5d13fa93e00..4f41d1d6abb 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/mod.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/mod.rs @@ -1,2 +1,2 @@ +pub(crate) mod auth; pub mod messages; -pub mod ocr; diff --git a/litellm-rust/crates/core/src/providers/azure_ai/ocr/mod.rs b/litellm-rust/crates/core/src/providers/azure_ai/ocr/mod.rs deleted file mode 100644 index f239b6921fa..00000000000 --- a/litellm-rust/crates/core/src/providers/azure_ai/ocr/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod transformation; diff --git a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs deleted file mode 100644 index 2af25a5639a..00000000000 --- a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs +++ /dev/null @@ -1,1381 +0,0 @@ -use std::collections::BTreeSet; - -use crate::error::{Error, json_type_name}; -use crate::ocr::transformation::{OcrAuthStrategy, OcrProviderConfig, OcrResponseHandling}; -use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData}; -use serde_json::{Map, Value, json}; - -use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; - -const AZURE_AI_API_KEY_ENV: &str = "AZURE_AI_API_KEY"; -const AZURE_AI_API_BASE_ENV: &str = "AZURE_AI_API_BASE"; -const AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_API_KEY"; -const AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV: &str = "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT"; -const AZURE_DOCUMENT_INTELLIGENCE_API_VERSION: &str = "2024-11-30"; -const AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI: i64 = 96; - -const AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS: &[&str] = - &["pages", "features", "req_format"]; - -pub struct AzureAiOcrConfig; -pub struct AzureDocumentIntelligenceOcrConfig; - -pub const AZURE_AI_OCR_CONFIG: AzureAiOcrConfig = AzureAiOcrConfig; -pub const AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG: AzureDocumentIntelligenceOcrConfig = - AzureDocumentIntelligenceOcrConfig; - -fn non_empty(value: Option<&str>) -> Option<&str> { - value.map(str::trim).filter(|value| !value.is_empty()) -} - -fn resolve_value( - explicit: Option<&str>, - env_name: &str, - env_lookup: &dyn Fn(&str) -> Option, - missing_message: &str, -) -> Result { - non_empty(explicit) - .map(str::to_string) - .or_else(|| env_lookup(env_name).filter(|value| !value.trim().is_empty())) - .ok_or_else(|| Error::Auth(missing_message.to_string())) -} - -pub fn resolve_azure_ai_api_key( - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - resolve_value( - api_key, - AZURE_AI_API_KEY_ENV, - env_lookup, - "Missing Azure AI API Key - A call is being made to Azure AI but no key is set either in the environment variables or via params", - ) -} - -pub fn resolve_azure_ai_api_base( - api_base: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - resolve_value( - api_base, - AZURE_AI_API_BASE_ENV, - env_lookup, - "Missing Azure AI API Base - Set AZURE_AI_API_BASE environment variable or pass api_base parameter", - ) -} - -pub fn complete_azure_ai_url( - api_base: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - let base = resolve_azure_ai_api_base(api_base, env_lookup)?; - Ok(format!( - "{}/providers/mistral/azure/ocr", - base.trim_end_matches('/') - )) -} - -pub fn resolve_document_intelligence_api_key( - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - resolve_value( - api_key, - AZURE_DOCUMENT_INTELLIGENCE_API_KEY_ENV, - env_lookup, - "Missing Azure Document Intelligence API Key - Set AZURE_DOCUMENT_INTELLIGENCE_API_KEY environment variable or pass api_key parameter", - ) -} - -pub fn resolve_document_intelligence_endpoint( - api_base: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - resolve_value( - api_base, - AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV, - env_lookup, - "Missing Azure Document Intelligence Endpoint - Set AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT environment variable or pass api_base parameter", - ) -} - -fn prepend_auth_header( - headers: Vec<(String, String)>, - name: &str, - value: String, -) -> Vec<(String, String)> { - std::iter::once((name.to_string(), value)) - .chain(headers) - .collect() -} - -pub fn validate_azure_ai_environment( - headers: Vec<(String, String)>, - api_key: Option<&str>, - azure_ad_token: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result, Error> { - if crate::http_utils::has_header(&headers, "Authorization") - || crate::http_utils::has_header(&headers, "Api-Key") - { - return Ok(headers); - } - if let Ok(api_key) = resolve_azure_ai_api_key(api_key, env_lookup) { - return Ok(prepend_auth_header(headers, "Api-Key", api_key)); - } - non_empty(azure_ad_token) - .map(|token| prepend_auth_header(headers, "Authorization", format!("Bearer {token}"))) - .ok_or_else(|| { - Error::Auth( - "Missing Azure AI credentials - set AZURE_AI_API_KEY or provide azure_ad_token" - .to_string(), - ) - }) -} - -pub fn validate_document_intelligence_environment( - headers: Vec<(String, String)>, - api_key: Option<&str>, - azure_ad_token: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result, Error> { - if crate::http_utils::has_header(&headers, "Authorization") - || crate::http_utils::has_header(&headers, "Ocp-Apim-Subscription-Key") - { - return Ok(headers); - } - if let Ok(api_key) = resolve_document_intelligence_api_key(api_key, env_lookup) { - return Ok(prepend_auth_header( - headers, - "Ocp-Apim-Subscription-Key", - api_key, - )); - } - non_empty(azure_ad_token) - .map(|token| prepend_auth_header(headers, "Authorization", format!("Bearer {token}"))) - .ok_or_else(|| { - Error::Auth( - "Missing Azure Document Intelligence credentials - set AZURE_DOCUMENT_INTELLIGENCE_API_KEY or provide azure_ad_token" - .to_string(), - ) - }) -} - -fn encode_model_id(model: &str) -> Result { - let model_id = model.rsplit('/').next().unwrap_or(model); - if matches!(model_id, "." | "..") { - return Err(Error::InvalidRequest( - "model_id cannot be a dot path segment".to_string(), - )); - } - Ok(model_id - .bytes() - .flat_map(|byte| match byte { - b'A'..=b'Z' | b'a'..=b'z' | b'0'..=b'9' | b'-' | b'_' | b'.' | b'~' => { - vec![byte as char] - } - _ => format!("%{byte:02X}").chars().collect(), - }) - .collect()) -} - -fn pages_token_is_valid(token: &str) -> bool { - let mut parts = token.split('-'); - let Some(start) = parts.next() else { - return false; - }; - if start.is_empty() || !start.chars().all(|ch| ch.is_ascii_digit()) { - return false; - } - match parts.next() { - None => true, - Some(end) => { - !end.is_empty() && end.chars().all(|ch| ch.is_ascii_digit()) && parts.next().is_none() - } - } -} - -fn normalize_pages_param(pages: &Value) -> Result, Error> { - match pages { - Value::String(value) => { - let normalized = value - .split(',') - .map(str::trim) - .collect::>() - .join(","); - if normalized.split(',').all(pages_token_is_valid) { - Ok(Some(normalized)) - } else { - Err(Error::InvalidRequest(format!( - "Invalid `pages` string for Azure Document Intelligence: {value:?}. Expected format like '1-3,5,7-9'." - ))) - } - } - Value::Array(values) => { - if values.is_empty() { - return Ok(None); - } - if values.iter().any(Value::is_boolean) { - return Err(Error::InvalidRequest( - "`pages` must be integers, not booleans".to_string(), - )); - } - if values.iter().all(Value::is_i64) { - let mut pages = BTreeSet::new(); - for value in values { - let page = value.as_i64().expect("checked is_i64"); - if page < 0 { - return Err(Error::InvalidRequest( - "`pages` integers must be >= 0 (Mistral 0-based indices)".to_string(), - )); - } - pages.insert(page + 1); - } - return Ok(Some( - pages - .into_iter() - .map(|page| page.to_string()) - .collect::>() - .join(","), - )); - } - if values.iter().all(Value::is_string) { - let normalized = values - .iter() - .filter_map(Value::as_str) - .map(str::trim) - .collect::>() - .join(","); - if normalized.split(',').all(pages_token_is_valid) { - return Ok(Some(normalized)); - } - return Err(Error::InvalidRequest(format!( - "Invalid `pages` list for Azure Document Intelligence: {values:?}. Expected tokens like '1' or '3-5'." - ))); - } - Err(Error::InvalidRequest( - "`pages` must be a list[int] (0-based, Mistral-style) or a string like '1-3,5,7-9'." - .to_string(), - )) - } - _ => Err(Error::InvalidRequest( - "`pages` must be a list[int] (0-based, Mistral-style) or a string like '1-3,5,7-9'." - .to_string(), - )), - } -} - -fn feature_token_is_valid(token: &str) -> bool { - let Some((first, rest)) = token.as_bytes().split_first() else { - return false; - }; - first.is_ascii_alphabetic() && rest.iter().all(u8::is_ascii_alphanumeric) -} - -fn invalid_features_error(features: &Value) -> Error { - Error::InvalidRequest(format!( - "Invalid `features` for Azure Document Intelligence: {features:?}. Expected a list of feature names or a comma-separated string like 'keyValuePairs' or 'keyValuePairs,languages'." - )) -} - -fn normalize_features_param(features: &Value) -> Result, Error> { - let normalized = match features { - Value::String(value) => value - .split(',') - .map(str::trim) - .collect::>() - .join(","), - Value::Array(values) if values.is_empty() => return Ok(None), - Value::Array(values) => values - .iter() - .map(Value::as_str) - .collect::>>() - .ok_or_else(|| invalid_features_error(features))? - .into_iter() - .map(str::trim) - .collect::>() - .join(","), - _ => return Err(invalid_features_error(features)), - }; - - if normalized.split(',').all(feature_token_is_valid) { - Ok(Some(normalized)) - } else { - Err(invalid_features_error(features)) - } -} - -fn normalize_req_format(req_format: &Value) -> Result { - match req_format.as_str() { - Some(value @ ("native" | "litellm")) => Ok(value.to_string()), - _ => Err(Error::InvalidRequest(format!( - "Invalid `req_format` for Azure Document Intelligence: {req_format:?}. Expected 'native' or 'litellm'." - ))), - } -} - -pub fn map_document_intelligence_ocr_params( - non_default_params: &Map, -) -> Result, Error> { - let mut mapped = Map::new(); - if let Some(pages) = non_default_params.get("pages") - && let Some(normalized) = normalize_pages_param(pages)? - { - mapped.insert("pages".to_string(), Value::String(normalized)); - } - if let Some(features) = non_default_params.get("features") - && let Some(normalized) = normalize_features_param(features)? - { - mapped.insert("features".to_string(), Value::String(normalized)); - } - if let Some(req_format) = non_default_params.get("req_format") { - mapped.insert( - "req_format".to_string(), - Value::String(normalize_req_format(req_format)?), - ); - } - Ok(mapped) -} - -pub fn complete_document_intelligence_url( - api_base: Option<&str>, - model: &str, - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - let endpoint = resolve_document_intelligence_endpoint(api_base, env_lookup)?; - let mut url = format!( - "{}/documentintelligence/documentModels/{}:analyze?api-version={}", - endpoint.trim_end_matches('/'), - encode_model_id(model)?, - AZURE_DOCUMENT_INTELLIGENCE_API_VERSION - ); - - if let Some(pages) = optional_params.get("pages") - && let Some(normalized) = normalize_pages_param(pages)? - { - url.push_str("&pages="); - url.push_str(&normalized); - } - - if let Some(features) = optional_params.get("features") - && let Some(normalized) = normalize_features_param(features)? - { - url.push_str("&features="); - url.push_str(&normalized); - } - - if let Some(req_format) = optional_params.get("req_format") { - normalize_req_format(req_format)?; - } - - Ok(url) -} - -fn document_url_from_mistral_document(document: &Value) -> Result<&str, Error> { - let object = document.as_object().ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(document), - })?; - let doc_type = object - .get("type") - .and_then(Value::as_str) - .ok_or(Error::MissingField("document.type"))?; - let field_name = match doc_type { - "document_url" => "document_url", - "image_url" => "image_url", - other => { - return Err(Error::InvalidRequest(format!( - "Invalid document type: {other}. Must be 'document_url' or 'image_url'" - ))); - } - }; - object - .get(field_name) - .and_then(Value::as_str) - .filter(|value| !value.is_empty()) - .ok_or(Error::MissingField(field_name)) -} - -fn extract_base64_from_data_uri(data_uri: &str) -> &str { - data_uri - .split_once(',') - .map(|(_, data)| data) - .unwrap_or(data_uri) -} - -fn page_markdown(page: &Map) -> String { - page.get("lines") - .and_then(Value::as_array) - .map(|lines| { - lines - .iter() - .filter_map(|line| line.get("content").and_then(Value::as_str)) - .collect::>() - .join("\n") - }) - .unwrap_or_default() -} - -fn page_dimensions(page: &Map) -> Value { - let width = page.get("width").and_then(Value::as_f64).unwrap_or(8.5); - let height = page.get("height").and_then(Value::as_f64).unwrap_or(11.0); - let unit = page.get("unit").and_then(Value::as_str).unwrap_or("inch"); - let (width, height) = if unit == "inch" { - ( - (width * AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI as f64) as i64, - (height * AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI as f64) as i64, - ) - } else { - (width as i64, height as i64) - }; - json!({ - "width": width, - "height": height, - "dpi": AZURE_DOCUMENT_INTELLIGENCE_DEFAULT_DPI, - }) -} - -fn transform_document_intelligence_response( - model: &str, - response_json: Value, - preserve_native_response: bool, -) -> Result { - let response = response_json - .as_object() - .ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(&response_json), - })?; - let status = response - .get("status") - .and_then(Value::as_str) - .ok_or(Error::MissingField("status"))?; - if status != "succeeded" { - return Err(Error::InvalidResponse(format!( - "Azure Document Intelligence analysis failed with status: {status}" - ))); - } - - let analyze_result = response.get("analyzeResult").and_then(Value::as_object); - let azure_pages = analyze_result - .and_then(|result| result.get("pages")) - .and_then(Value::as_array) - .cloned() - .unwrap_or_default(); - let pages = azure_pages - .iter() - .filter_map(Value::as_object) - .map(|page| { - let page_number = page.get("pageNumber").and_then(Value::as_i64).unwrap_or(1); - json!({ - "index": page_number - 1, - "markdown": page_markdown(page), - "dimensions": page_dimensions(page), - }) - }) - .collect::>(); - let extra_fields = ["content", "tables", "keyValuePairs"] - .into_iter() - .map(|field| { - ( - field.to_string(), - analyze_result - .and_then(|result| result.get(field)) - .cloned() - .unwrap_or(Value::Null), - ) - }) - .collect(); - - Ok(LiteLLMOcrResponse { - usage_info: Some(json!({ - "pages_processed": pages.len(), - "doc_size_bytes": null, - })), - pages, - model: model.to_string(), - document_annotation: None, - object: "ocr".to_string(), - extra_fields, - provider_native_response: preserve_native_response.then_some(response_json), - }) -} - -impl OcrProviderConfig for AzureAiOcrConfig { - fn supported_ocr_params(&self) -> &'static [&'static str] { - MISTRAL_OCR_CONFIG.supported_ocr_params() - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn transform_ocr_request( - &self, - model: &str, - document: Value, - optional_params: Map, - ) -> Result { - MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params) - } - - fn transform_ocr_response( - &self, - model: &str, - response_json: Value, - ) -> Result { - MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn complete_url( - &self, - api_base: Option<&str>, - _model: &str, - _optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - complete_azure_ai_url(api_base, env_lookup) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_azure_ai_api_key(api_key, env_lookup) - } - - fn requires_data_uri_document(&self) -> bool { - true - } -} - -impl OcrProviderConfig for AzureDocumentIntelligenceOcrConfig { - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn supported_ocr_params(&self) -> &'static [&'static str] { - AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn map_ocr_params(&self, non_default_params: &Map) -> Map { - map_document_intelligence_ocr_params(non_default_params).unwrap_or_else(|_| { - non_default_params - .iter() - .filter(|(name, _)| { - AZURE_DOCUMENT_INTELLIGENCE_SUPPORTED_OCR_PARAMS.contains(&name.as_str()) - }) - .map(|(name, value)| (name.clone(), value.clone())) - .collect() - }) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn transform_ocr_request( - &self, - _model: &str, - document: Value, - _optional_params: Map, - ) -> Result { - let document_url = document_url_from_mistral_document(&document)?; - let mut data = Map::new(); - if document_url.starts_with("data:") { - data.insert( - "base64Source".to_string(), - Value::String(extract_base64_from_data_uri(document_url).to_string()), - ); - } else { - data.insert( - "urlSource".to_string(), - Value::String(document_url.to_string()), - ); - } - Ok(OcrRequestData { - data: Value::Object(data), - files: None, - }) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn transform_ocr_response( - &self, - model: &str, - response_json: Value, - ) -> Result { - transform_document_intelligence_response(model, response_json, false) - } - - fn transform_ocr_response_with_params( - &self, - model: &str, - response_json: Value, - optional_params: &Map, - ) -> Result { - transform_document_intelligence_response( - model, - response_json, - optional_params.get("req_format").and_then(Value::as_str) == Some("native"), - ) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn complete_url( - &self, - api_base: Option<&str>, - model: &str, - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - complete_document_intelligence_url(api_base, model, optional_params, env_lookup) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_document_intelligence_api_key(api_key, env_lookup) - } - - fn auth_strategy(&self) -> OcrAuthStrategy { - OcrAuthStrategy::Header("Ocp-Apim-Subscription-Key") - } - - fn response_handling(&self) -> OcrResponseHandling { - OcrResponseHandling::AzureDocumentIntelligencePoll - } -} - -#[cfg(test)] -mod tests { - use super::*; - use rstest::{fixture, rstest}; - - const ENDPOINT: &str = "https://example.cognitiveservices.azure.com"; - - #[fixture] - fn document_intelligence_config() -> AzureDocumentIntelligenceOcrConfig { - AzureDocumentIntelligenceOcrConfig - } - - fn header_value<'a>(headers: &'a [(String, String)], name: &str) -> Option<&'a str> { - headers - .iter() - .find(|(header_name, _)| header_name.eq_ignore_ascii_case(name)) - .map(|(_, value)| value.as_str()) - } - - #[fixture] - fn native_operation() -> Value { - json!({ - "status": "succeeded", - "createdDateTime": "2026-07-02T00:00:00Z", - "lastUpdatedDateTime": "2026-07-02T00:00:05Z", - "analyzeResult": { - "content": "Invoice\nInvoice No: INV-12345\nTotal: $100.00", - "pages": [{ - "pageNumber": 1, - "width": 8.5, - "height": 11, - "unit": "inch", - "angle": 0.13, - "lines": [ - {"content": "Invoice"}, - {"content": "Invoice No: INV-12345"}, - {"content": "Total: $100.00"} - ], - "words": [{"content": "Invoice", "confidence": 0.994}] - }], - "tables": [ - { - "rowCount": 2, - "columnCount": 2, - "cells": [ - {"kind": "columnHeader", "rowIndex": 0, "columnIndex": 0, "content": "Item"}, - {"kind": "columnHeader", "rowIndex": 0, "columnIndex": 1, "content": "Price"}, - {"rowIndex": 1, "columnIndex": 0, "content": "Widget"}, - {"rowIndex": 1, "columnIndex": 1, "content": "$100.00"} - ] - }, - { - "rowCount": 1, - "columnCount": 1, - "cells": [{"rowIndex": 0, "columnIndex": 0, "content": "Totals"}] - } - ], - "keyValuePairs": [ - { - "key": {"content": "Invoice No"}, - "value": {"content": "INV-12345"}, - "confidence": 0.98 - }, - { - "key": {"content": "Total"}, - "value": {"content": "$100.00"}, - "confidence": 0.95 - } - ], - "paragraphs": [{"content": "Invoice"}] - } - }) - } - - fn assert_native_fields_preserved(response: &LiteLLMOcrResponse, operation: &Value) { - let analyze_result = &operation["analyzeResult"]; - - assert_eq!(response.extra_fields["content"], analyze_result["content"]); - assert_eq!(response.extra_fields["tables"], analyze_result["tables"]); - assert_eq!( - response.extra_fields["keyValuePairs"], - analyze_result["keyValuePairs"] - ); - assert_eq!(response.object, "ocr"); - assert_eq!( - response.usage_info, - Some(json!({"pages_processed": 1, "doc_size_bytes": null})) - ); - assert_eq!(response.pages[0]["index"], 0); - assert_eq!( - response.pages[0]["markdown"], - "Invoice\nInvoice No: INV-12345\nTotal: $100.00" - ); - assert_eq!( - response.pages[0]["dimensions"], - json!({"width": 816, "height": 1056, "dpi": 96}) - ); - } - - #[test] - fn azure_ai_reuses_mistral_body_transform() { - let body = AZURE_AI_OCR_CONFIG - .transform_ocr_request( - "pixtral-12b-2409", - json!({"type": "document_url", "document_url": "data:application/pdf;base64,abc"}), - serde_json::Map::from_iter([("include_image_base64".to_string(), json!(true))]), - ) - .expect("request transforms") - .data; - - assert_eq!(body["model"], "pixtral-12b-2409"); - assert_eq!(body["include_image_base64"], true); - assert_eq!( - body["document"]["document_url"], - "data:application/pdf;base64,abc" - ); - } - - #[test] - fn document_intelligence_url_normalizes_zero_based_pages() { - let params = serde_json::Map::from_iter([("pages".to_string(), json!([2, 0, 2]))]); - let url = complete_document_intelligence_url( - Some("https://example.cognitiveservices.azure.com/"), - "azure_ai/doc-intelligence/prebuilt-layout", - ¶ms, - &|_| None, - ) - .expect("url builds"); - - assert_eq!( - url, - "https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&pages=1,3" - ); - } - - #[test] - fn document_intelligence_url_normalizes_features() { - let params = serde_json::Map::from_iter([( - "features".to_string(), - json!("keyValuePairs, languages"), - )]); - let url = complete_document_intelligence_url( - Some("https://example.cognitiveservices.azure.com"), - "prebuilt-layout", - ¶ms, - &|_| None, - ) - .expect("url builds"); - - assert_eq!( - url, - "https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&features=keyValuePairs,languages" - ); - } - - #[test] - fn document_intelligence_url_combines_pages_and_feature_list() { - let params = serde_json::Map::from_iter([ - ("pages".to_string(), json!([0, 1, 2])), - ( - "features".to_string(), - json!([" keyValuePairs ", "languages"]), - ), - ]); - let url = complete_document_intelligence_url( - Some("https://example.cognitiveservices.azure.com"), - "prebuilt-layout", - ¶ms, - &|_| None, - ) - .expect("url builds"); - - assert_eq!( - url, - "https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30&pages=1,2,3&features=keyValuePairs,languages" - ); - } - - #[test] - fn document_intelligence_url_omits_empty_feature_list() { - let params = serde_json::Map::from_iter([("features".to_string(), json!([]))]); - assert!( - map_document_intelligence_ocr_params(¶ms) - .expect("empty features map") - .is_empty() - ); - let url = complete_document_intelligence_url( - Some("https://example.cognitiveservices.azure.com"), - "prebuilt-layout", - ¶ms, - &|_| None, - ) - .expect("url builds"); - - assert_eq!( - url, - "https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout:analyze?api-version=2024-11-30" - ); - } - - #[rstest] - #[case::query_injection(json!("keyValuePairs&pages=9"))] - #[case::spaces(json!("key value pairs"))] - #[case::empty_string(json!(""))] - #[case::integer_list(json!([1, 2]))] - #[case::nested_list(json!([["keyValuePairs"]]))] - #[case::object(json!({"feature": "keyValuePairs"}))] - #[case::number(json!(5))] - fn document_intelligence_mapping_rejects_invalid_features(#[case] features: Value) { - let params = serde_json::Map::from_iter([("features".to_string(), features)]); - let error = - map_document_intelligence_ocr_params(¶ms).expect_err("invalid features must fail"); - - assert!(matches!( - error, - Error::InvalidRequest(message) if message.contains("Invalid `features`") - )); - } - - #[rstest] - #[case::single_list(json!(["keyValuePairs"]), "keyValuePairs")] - #[case::multiple_list( - json!(["keyValuePairs", "languages"]), - "keyValuePairs,languages" - )] - #[case::single_string(json!("keyValuePairs"), "keyValuePairs")] - #[case::comma_separated(json!("keyValuePairs,languages"), "keyValuePairs,languages")] - #[case::spaces(json!("keyValuePairs, languages"), "keyValuePairs,languages")] - fn document_intelligence_maps_features(#[case] features: Value, #[case] expected: &str) { - let params = Map::from_iter([ - ("features".to_string(), features), - ("unsupported".to_string(), json!(true)), - ]); - - assert_eq!( - AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG.map_ocr_params(¶ms), - Map::from_iter([("features".to_string(), json!(expected))]) - ); - } - - #[test] - fn document_intelligence_request_uses_base64_source_for_data_uri() { - let body = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_request( - "prebuilt-read", - json!({"type": "document_url", "document_url": "data:application/pdf;base64,abc123"}), - Map::new(), - ) - .expect("request transforms") - .data; - - assert_eq!(body, json!({"base64Source": "abc123"})); - } - - #[rstest] - fn document_intelligence_response_normalizes_pages(native_operation: Value) { - let response = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_response("prebuilt-layout", native_operation.clone()) - .expect("response transforms"); - - assert_native_fields_preserved(&response, &native_operation); - } - - #[test] - fn azure_document_intelligence_model_id_is_encoded() { - let url = complete_document_intelligence_url( - Some(ENDPOINT), - "prebuilt-layout?x=1#frag", - &Map::new(), - &|_| None, - ) - .expect("url builds"); - - assert_eq!( - url, - "https://example.cognitiveservices.azure.com/documentintelligence/documentModels/prebuilt-layout%3Fx%3D1%23frag:analyze?api-version=2024-11-30" - ); - } - - #[test] - fn azure_document_intelligence_dot_segment_model_id_is_rejected() { - let error = complete_document_intelligence_url( - Some(ENDPOINT), - "azure_ai/doc-intelligence/..", - &Map::new(), - &|_| None, - ) - .expect_err("dot segment must fail"); - - assert_eq!( - error, - Error::InvalidRequest("model_id cannot be a dot path segment".to_string()) - ); - } - - #[rstest] - fn document_intelligence_async_response_preserves_normalized_fields(native_operation: Value) { - let response = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_response( - "azure_ai/doc-intelligence/prebuilt-layout", - native_operation.clone(), - ) - .expect("response transforms"); - - assert_native_fields_preserved(&response, &native_operation); - } - - #[test] - fn document_intelligence_response_tolerates_missing_native_fields() { - let response = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_response( - "azure_ai/doc-intelligence/prebuilt-read", - json!({ - "status": "succeeded", - "analyzeResult": { - "pages": [{ - "pageNumber": 1, - "width": 8.5, - "height": 11, - "unit": "inch", - "lines": [{"content": "hello"}] - }] - } - }), - ) - .expect("missing optional fields are allowed"); - - assert_eq!(response.pages[0]["markdown"], "hello"); - assert_eq!(response.extra_fields["content"], Value::Null); - assert_eq!(response.extra_fields["tables"], Value::Null); - assert_eq!(response.extra_fields["keyValuePairs"], Value::Null); - } - - #[test] - fn document_intelligence_non_succeeded_status_is_rejected() { - let error = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_response( - "azure_ai/doc-intelligence/prebuilt-layout", - json!({"status": "failed"}), - ) - .expect_err("failed status must fail"); - - assert_eq!( - error, - Error::InvalidResponse( - "Azure Document Intelligence analysis failed with status: failed".to_string() - ) - ); - } - - #[test] - fn document_intelligence_supported_params_include_features() { - assert_eq!( - AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG.supported_ocr_params(), - &["pages", "features", "req_format"] - ); - } - - #[rstest] - fn document_intelligence_native_format_carries_raw_operation(native_operation: Value) { - let response = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_response_with_params( - "azure_ai/doc-intelligence/prebuilt-layout", - native_operation.clone(), - &Map::from_iter([("req_format".to_string(), json!("native"))]), - ) - .expect("native response transforms"); - - assert_eq!( - response.provider_native_response, - Some(native_operation.clone()) - ); - assert_native_fields_preserved(&response, &native_operation); - } - - #[rstest] - fn document_intelligence_async_native_format_carries_raw_operation(native_operation: Value) { - let response = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_response_with_params( - "azure_ai/doc-intelligence/prebuilt-layout", - native_operation.clone(), - &Map::from_iter([("req_format".to_string(), json!("native"))]), - ) - .expect("native response transforms"); - - assert_eq!( - response.provider_native_response, - Some(native_operation.clone()) - ); - assert_native_fields_preserved(&response, &native_operation); - } - - #[rstest] - #[case::default(Map::new())] - #[case::litellm(Map::from_iter([("req_format".to_string(), json!("litellm"))]))] - fn document_intelligence_default_format_omits_raw_operation( - #[case] optional_params: Map, - native_operation: Value, - ) { - let response = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_response_with_params( - "azure_ai/doc-intelligence/prebuilt-layout", - native_operation.clone(), - &optional_params, - ) - .expect("response transforms"); - - assert_eq!(response.provider_native_response, None); - assert_native_fields_preserved(&response, &native_operation); - } - - #[rstest] - #[case::native("native")] - #[case::litellm("litellm")] - fn document_intelligence_maps_req_format(#[case] req_format: &str) { - let mapped = map_document_intelligence_ocr_params(&Map::from_iter([( - "req_format".to_string(), - json!(req_format), - )])) - .expect("req_format maps"); - - assert_eq!( - mapped, - Map::from_iter([("req_format".to_string(), json!(req_format))]) - ); - } - - #[test] - fn document_intelligence_rejects_unknown_req_format() { - let error = map_document_intelligence_ocr_params(&Map::from_iter([( - "req_format".to_string(), - json!("azure"), - )])) - .expect_err("unknown req_format must fail"); - - assert!( - matches!(error, Error::InvalidRequest(message) if message.contains("Invalid `req_format`")) - ); - } - - #[test] - fn document_intelligence_url_omits_req_format() { - let url = complete_document_intelligence_url( - Some(ENDPOINT), - "prebuilt-layout", - &Map::from_iter([("req_format".to_string(), json!("native"))]), - &|_| None, - ) - .expect("url builds"); - - assert!(!url.contains("req_format")); - } - - #[test] - fn document_intelligence_validate_environment_uses_subscription_key() { - let headers = - validate_document_intelligence_environment(Vec::new(), Some("my-key"), None, &|_| None) - .expect("api key authenticates"); - - assert_eq!( - header_value(&headers, "Ocp-Apim-Subscription-Key"), - Some("my-key") - ); - } - - #[test] - fn document_intelligence_validate_environment_falls_back_to_entra_token() { - let headers = validate_document_intelligence_environment( - Vec::new(), - None, - Some("entra-token"), - &|_| None, - ) - .expect("Entra token authenticates"); - - assert_eq!( - header_value(&headers, "Authorization"), - Some("Bearer entra-token") - ); - assert_eq!(header_value(&headers, "Ocp-Apim-Subscription-Key"), None); - } - - #[test] - fn document_intelligence_supported_params_include_pages_features_and_req_format() { - assert_eq!( - AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG.supported_ocr_params(), - &["pages", "features", "req_format"] - ); - } - - #[test] - fn document_intelligence_maps_zero_based_page_list() { - let mapped = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!([0, 1, 2]), - )])) - .expect("pages map"); - - assert_eq!( - mapped, - Map::from_iter([("pages".to_string(), json!("1,2,3"))]) - ); - } - - #[test] - fn document_intelligence_page_mapping_dedupes_and_sorts() { - let mapped = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!([2, 0, 0, 1]), - )])) - .expect("pages map"); - - assert_eq!(mapped["pages"], "1,2,3"); - } - - #[test] - fn document_intelligence_page_mapping_omits_empty_list() { - let mapped = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!([]), - )])) - .expect("empty pages map"); - - assert!(mapped.is_empty()); - } - - #[test] - fn document_intelligence_page_mapping_accepts_native_range() { - let mapped = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!("3-9"), - )])) - .expect("range maps"); - - assert_eq!(mapped["pages"], "3-9"); - } - - #[test] - fn document_intelligence_page_mapping_strips_spaces() { - let mapped = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!("1-3, 5"), - )])) - .expect("range maps"); - - assert_eq!(mapped["pages"], "1-3,5"); - } - - #[test] - fn document_intelligence_page_mapping_accepts_string_tokens() { - let mapped = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!(["1", "3-5"]), - )])) - .expect("tokens map"); - - assert_eq!(mapped["pages"], "1,3-5"); - } - - #[test] - fn document_intelligence_page_mapping_rejects_invalid_string() { - let error = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!("a,b"), - )])) - .expect_err("invalid pages must fail"); - - assert!( - matches!(error, Error::InvalidRequest(message) if message.contains("Invalid `pages` string")) - ); - } - - #[test] - fn document_intelligence_page_mapping_rejects_negative_index() { - let error = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!([-1]), - )])) - .expect_err("negative pages must fail"); - - assert!( - matches!(error, Error::InvalidRequest(message) if message.contains("must be >= 0")) - ); - } - - #[test] - fn document_intelligence_page_mapping_rejects_bool_list() { - let error = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!([true, false]), - )])) - .expect_err("boolean pages must fail"); - - assert!( - matches!(error, Error::InvalidRequest(message) if message.contains("integers, not booleans")) - ); - } - - #[test] - fn document_intelligence_page_mapping_rejects_unsupported_type() { - let error = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!(5), - )])) - .expect_err("unsupported pages must fail"); - - assert!( - matches!(error, Error::InvalidRequest(message) if message.contains("Mistral-style")) - ); - } - - #[test] - fn document_intelligence_url_appends_pages_query() { - let url = complete_document_intelligence_url( - Some("https://example.cognitiveservices.azure.com/"), - "azure_ai/doc-intelligence/prebuilt-layout", - &Map::from_iter([("pages".to_string(), json!("1-3,5"))]), - &|_| None, - ) - .expect("url builds"); - - assert!(url.contains("api-version=2024-11-30")); - assert!(url.contains("pages=1-3,5")); - assert!(url.contains("/documentintelligence/documentModels/prebuilt-layout:analyze")); - } - - #[test] - fn document_intelligence_url_has_no_pages_when_params_are_empty() { - let url = complete_document_intelligence_url( - Some(ENDPOINT), - "prebuilt-layout", - &Map::new(), - &|_| None, - ) - .expect("url builds"); - - assert!(!url.contains("pages=")); - } - - #[rstest] - fn document_intelligence_request_keeps_pages_out_of_body( - document_intelligence_config: AzureDocumentIntelligenceOcrConfig, - ) { - let request = document_intelligence_config - .transform_ocr_request( - "prebuilt-layout", - json!({"type": "document_url", "document_url": "https://example.com/x.pdf"}), - Map::from_iter([("pages".to_string(), json!("1,2,3"))]), - ) - .expect("request transforms"); - - assert_eq!( - request.data, - json!({"urlSource": "https://example.com/x.pdf"}) - ); - } - - #[test] - fn document_intelligence_mistral_pages_flow_to_query_only() { - let mapped = map_document_intelligence_ocr_params(&Map::from_iter([( - "pages".to_string(), - json!([2, 3, 4, 5, 6, 7, 8]), - )])) - .expect("pages map"); - let url = - complete_document_intelligence_url(Some(ENDPOINT), "prebuilt-layout", &mapped, &|_| { - None - }) - .expect("url builds"); - let request = AZURE_DOCUMENT_INTELLIGENCE_OCR_CONFIG - .transform_ocr_request( - "prebuilt-layout", - json!({"type": "document_url", "document_url": "https://example.com/x.pdf"}), - mapped, - ) - .expect("request transforms"); - - assert!(url.contains("pages=3,4,5,6,7,8,9")); - assert_eq!( - request.data, - json!({"urlSource": "https://example.com/x.pdf"}) - ); - } - - #[test] - fn document_intelligence_endpoint_ignores_generic_azure_ai_base() { - let resolved = resolve_document_intelligence_endpoint(None, &|name| match name { - AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV => Some(ENDPOINT.to_string()), - AZURE_AI_API_BASE_ENV => Some("https://generic.example.com".to_string()), - _ => None, - }) - .expect("endpoint resolves"); - - assert_eq!(resolved, ENDPOINT); - } - - #[test] - fn document_intelligence_endpoint_honors_explicit_api_base() { - let resolved = resolve_document_intelligence_endpoint( - Some("https://my-di.cognitiveservices.azure.com"), - &|name| match name { - AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV => Some(ENDPOINT.to_string()), - AZURE_AI_API_BASE_ENV => Some("https://generic.example.com".to_string()), - _ => None, - }, - ) - .expect("endpoint resolves"); - - assert_eq!(resolved, "https://my-di.cognitiveservices.azure.com"); - } - - #[test] - fn azure_ai_mistral_ocr_uses_generic_api_base() { - let resolved = resolve_azure_ai_api_base(None, &|name| match name { - AZURE_AI_API_BASE_ENV => Some("https://generic-azure-ai.example.com".to_string()), - AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT_ENV => Some(ENDPOINT.to_string()), - _ => None, - }) - .expect("api base resolves"); - - assert_eq!(resolved, "https://generic-azure-ai.example.com"); - } - - #[test] - fn azure_ai_ocr_authenticates_with_entra_token() { - let headers = - validate_azure_ai_environment(Vec::new(), None, Some("entra-token"), &|_| None) - .expect("Entra token authenticates"); - - assert_eq!( - header_value(&headers, "Authorization"), - Some("Bearer entra-token") - ); - } -} diff --git a/litellm-rust/crates/core/src/providers/mistral/mod.rs b/litellm-rust/crates/core/src/providers/mistral/mod.rs deleted file mode 100644 index 3621ff6a2fd..00000000000 --- a/litellm-rust/crates/core/src/providers/mistral/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod ocr; diff --git a/litellm-rust/crates/core/src/providers/mistral/ocr/mod.rs b/litellm-rust/crates/core/src/providers/mistral/ocr/mod.rs deleted file mode 100644 index f239b6921fa..00000000000 --- a/litellm-rust/crates/core/src/providers/mistral/ocr/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod transformation; diff --git a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs deleted file mode 100644 index 044fc587c22..00000000000 --- a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs +++ /dev/null @@ -1,436 +0,0 @@ -use crate::error::{Error, json_type_name}; -use crate::ocr::transformation::OcrProviderConfig; -use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData}; -use serde_json::{Map, Value}; - -const SUPPORTED_OCR_PARAMS: &[&str] = &[ - "pages", - "include_image_base64", - "image_limit", - "image_min_size", - "bbox_annotation_format", - "document_annotation_format", - "document_annotation_prompt", - "extract_header", - "extract_footer", - "table_format", - "confidence_scores_granularity", - "include_blocks", - "id", -]; - -/// Default Mistral API base, used when the caller does not override `api_base`. -pub const MISTRAL_DEFAULT_API_BASE: &str = "https://api.mistral.ai/v1"; - -/// Environment variable holding the Mistral API key. -pub const MISTRAL_API_KEY_ENV: &str = "MISTRAL_API_KEY"; - -/// Error message raised when no Mistral API key can be resolved. -pub const MISSING_KEY_MESSAGE: &str = "Missing Mistral API Key - A call is being made to Mistral but no key is set either in the environment variables or via params"; - -/// Build the complete OCR endpoint URL, de-duplicating a trailing `/v1`. -/// -/// Blank/whitespace `api_base` is treated as absent (guard at resolution time). -pub fn complete_url(api_base: Option<&str>) -> String { - let base = api_base - .map(str::trim) - .filter(|base| !base.is_empty()) - .unwrap_or(MISTRAL_DEFAULT_API_BASE) - .trim_end_matches('/'); - - if base.ends_with("/v1") { - format!("{base}/ocr") - } else { - format!("{base}/v1/ocr") - } -} - -/// Resolve the Mistral API key from the explicit param or the environment. -/// -/// Blank/whitespace values are treated as absent. Returns `Error::Auth` -/// when no usable key is available. -/// -/// Note: the env fallback only reads the process environment. Secret-manager -/// backends (AWS/Azure/GCP/Vault) are resolved on the Python side and passed in -/// via `api_key`; this fallback is a last resort for direct/standalone use. -pub fn resolve_api_key( - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - api_key - .map(str::trim) - .filter(|key| !key.is_empty()) - .map(str::to_string) - .or_else(|| env_lookup(MISTRAL_API_KEY_ENV).filter(|key| !key.trim().is_empty())) - .ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string())) -} - -pub struct MistralOcrConfig; - -pub const MISTRAL_OCR_CONFIG: MistralOcrConfig = MistralOcrConfig; - -impl OcrProviderConfig for MistralOcrConfig { - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn supported_ocr_params(&self) -> &'static [&'static str] { - SUPPORTED_OCR_PARAMS - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn transform_ocr_request( - &self, - model: &str, - document: Value, - optional_params: Map, - ) -> Result { - if !document.is_object() { - return Err(Error::InvalidType { - expected: "object", - actual: json_type_name(&document), - }); - } - - let mut data = Map::new(); - data.insert("model".to_string(), Value::String(model.to_string())); - data.insert("document".to_string(), document); - for (param, value) in optional_params { - data.insert(param, value); - } - - Ok(OcrRequestData { - data: Value::Object(data), - files: None, - }) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn transform_ocr_response( - &self, - model: &str, - response_json: Value, - ) -> Result { - let response_object = response_json - .as_object() - .ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(&response_json), - })?; - - let pages = response_object - .get("pages") - .and_then(Value::as_array) - .cloned() - .unwrap_or_default(); - let model = response_object - .get("model") - .and_then(Value::as_str) - .unwrap_or(model) - .to_string(); - let document_annotation = response_object.get("document_annotation").cloned(); - let usage_info = response_object.get("usage_info").cloned(); - - Ok(LiteLLMOcrResponse { - pages, - model, - document_annotation, - usage_info, - object: "ocr".to_string(), - extra_fields: Map::new(), - provider_native_response: None, - }) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn complete_url( - &self, - api_base: Option<&str>, - _model: &str, - _optional_params: &Map, - _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - Ok(complete_url(api_base)) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_api_key(api_key, env_lookup) - } -} - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub fn supported_ocr_params() -> &'static [&'static str] { - MISTRAL_OCR_CONFIG.supported_ocr_params() -} - -pub fn map_ocr_params(non_default_params: &Map) -> Map { - MISTRAL_OCR_CONFIG.map_ocr_params(non_default_params) -} - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub fn transform_ocr_request( - model: &str, - document: Value, - optional_params: Map, -) -> Result { - MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params) -} - -#[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] -pub fn transform_ocr_response( - model: &str, - response_json: Value, -) -> Result { - MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json) -} - -#[cfg(test)] -mod tests { - use super::*; - use serde_json::json; - - #[test] - fn extract_header_is_a_supported_ocr_param() { - assert!(supported_ocr_params().contains(&"extract_header")); - } - - #[test] - fn extract_footer_is_a_supported_ocr_param() { - assert!(supported_ocr_params().contains(&"extract_footer")); - } - - #[test] - fn existing_ocr_params_remain_supported() { - for param in [ - "pages", - "include_image_base64", - "image_limit", - "image_min_size", - "bbox_annotation_format", - "document_annotation_format", - ] { - assert!(supported_ocr_params().contains(¶m)); - } - } - - #[test] - fn map_ocr_params_forwards_extract_header() { - let params = json!({"extract_header": true}); - assert_eq!( - map_ocr_params(params.as_object().unwrap()), - params.as_object().unwrap().clone() - ); - } - - #[test] - fn map_ocr_params_forwards_extract_footer() { - let params = json!({"extract_footer": true}); - assert_eq!( - map_ocr_params(params.as_object().unwrap()), - params.as_object().unwrap().clone() - ); - } - - #[test] - fn map_ocr_params_forwards_extract_header_and_footer() { - let params = json!({"extract_header": true, "extract_footer": false}); - assert_eq!( - map_ocr_params(params.as_object().unwrap()), - params.as_object().unwrap().clone() - ); - } - - #[test] - fn map_ocr_params_drops_unknown_params() { - let params = json!({"extract_header": true, "unsupported_param": "value"}); - let mapped = map_ocr_params(params.as_object().unwrap()); - assert_eq!(mapped.get("extract_header"), Some(&json!(true))); - assert!(!mapped.contains_key("unsupported_param")); - } - - #[test] - fn new_ocr_params_are_supported() { - for param in [ - "table_format", - "confidence_scores_granularity", - "document_annotation_prompt", - "include_blocks", - "id", - ] { - assert!(supported_ocr_params().contains(¶m)); - } - } - - #[test] - fn map_ocr_params_forwards_new_ocr_params() { - for (param, value) in [ - ("table_format", json!("html")), - ("confidence_scores_granularity", json!("word")), - ( - "document_annotation_prompt", - json!("Extract all invoice line items"), - ), - ("include_blocks", json!(true)), - ("id", json!("req-123")), - ] { - let params = json!({param: value}); - assert_eq!( - map_ocr_params(params.as_object().unwrap()), - params.as_object().unwrap().clone() - ); - } - } - - #[test] - fn transform_ocr_request_includes_each_optional_param() { - let document = json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }); - for (param, value) in [ - ("table_format", json!("html")), - ("confidence_scores_granularity", json!("word")), - ( - "document_annotation_prompt", - json!("Extract all invoice line items"), - ), - ("id", json!("req-123")), - ("extract_header", json!(true)), - ("include_blocks", json!(true)), - ("pages", json!([0, 1])), - ] { - let result = transform_ocr_request( - "mistral-ocr-latest", - document.clone(), - json!({param: value}).as_object().unwrap().clone(), - ) - .expect("request should transform"); - assert_eq!(result.data.get(param), Some(&value)); - assert_eq!(result.data.get("model"), Some(&json!("mistral-ocr-latest"))); - assert_eq!(result.data.get("document"), Some(&document)); - assert_eq!(result.files, None); - } - } - - #[test] - fn transform_ocr_request_includes_multiple_new_params() { - let document = json!({ - "type": "document_url", - "document_url": "https://example.com/doc.pdf" - }); - let optional_params = json!({ - "table_format": "html", - "confidence_scores_granularity": "page", - "extract_header": true - }) - .as_object() - .unwrap() - .clone(); - let result = transform_ocr_request("mistral-ocr-latest", document, optional_params) - .expect("request should transform"); - assert_eq!(result.data.get("table_format"), Some(&json!("html"))); - assert_eq!( - result.data.get("confidence_scores_granularity"), - Some(&json!("page")) - ); - assert_eq!(result.data.get("extract_header"), Some(&json!(true))); - } - - #[test] - fn transform_ocr_response_preserves_blocks_and_confidence_scores() { - let blocks = json!([{"type": "title", "content": "Invoice"}]); - let confidence_scores = json!({"page": 0.98}); - let response = json!({ - "pages": [{"index": 0, "markdown": "# Invoice", "blocks": blocks, "confidence_scores": confidence_scores}], - "model": "mistral-ocr-4-0", - "usage_info": {"pages_processed": 1} - }); - let result = - transform_ocr_response("mistral-ocr-4-0", response).expect("response should transform"); - assert_eq!(result.pages[0].get("blocks"), Some(&blocks)); - assert_eq!( - result.pages[0].get("confidence_scores"), - Some(&confidence_scores) - ); - } - - #[test] - fn transform_ocr_response_preserves_ocr4_page_fields() { - let response = json!({ - "pages": [{"index": 0, "markdown": "table page", "tables": [{"rows": 2, "cols": 3}], "hyperlinks": ["https://example.com"], "header": "Acme Corp", "footer": "Page 1"}], - "model": "mistral-ocr-4-0", - "usage_info": {"pages_processed": 1} - }); - let result = transform_ocr_response("mistral-ocr-4-0", response.clone()) - .expect("response should transform"); - assert_eq!(result.pages[0], response["pages"][0]); - } - - #[test] - fn transform_ocr_request_rejects_non_object_document() { - let err = transform_ocr_request("mistral-ocr-latest", json!("bad"), Map::new()) - .expect_err("string document should be rejected"); - - assert_eq!( - err, - Error::InvalidType { - expected: "object", - actual: "string", - } - ); - } - - #[test] - fn transform_ocr_response_normalizes_mistral_json() { - let response = json!({ - "pages": [{"index": 0, "markdown": "hello"}], - "model": "mistral-ocr-2505-completion", - "document_annotation": null, - "usage_info": {"pages_processed": 1} - }); - - let result = transform_ocr_response("mistral-ocr-latest", response) - .expect("response should transform"); - - assert_eq!(result.pages, vec![json!({"index": 0, "markdown": "hello"})]); - assert_eq!(result.model, "mistral-ocr-2505-completion"); - assert_eq!(result.document_annotation, Some(Value::Null)); - assert_eq!(result.usage_info, Some(json!({"pages_processed": 1}))); - assert_eq!(result.object, "ocr"); - } - - #[test] - fn complete_url_defaults_and_dedupes_v1() { - assert_eq!(complete_url(None), "https://api.mistral.ai/v1/ocr"); - assert_eq!(complete_url(Some(" ")), "https://api.mistral.ai/v1/ocr"); - assert_eq!( - complete_url(Some("https://proxy.internal")), - "https://proxy.internal/v1/ocr" - ); - assert_eq!( - complete_url(Some("https://proxy.internal/v1/")), - "https://proxy.internal/v1/ocr" - ); - } - - #[test] - fn resolve_api_key_prefers_param_then_env() { - let no_env = |_: &str| None; - assert_eq!( - resolve_api_key(Some("sk-param"), &no_env).unwrap(), - "sk-param" - ); - - let with_env = |key: &str| (key == MISTRAL_API_KEY_ENV).then(|| "sk-env".to_string()); - assert_eq!(resolve_api_key(None, &with_env).unwrap(), "sk-env"); - // Blank param falls through to the environment. - assert_eq!(resolve_api_key(Some(" "), &with_env).unwrap(), "sk-env"); - } - - #[test] - fn resolve_api_key_errors_when_absent() { - let err = resolve_api_key(None, &|_| None).expect_err("missing key should error"); - assert_eq!(err, Error::Auth(MISSING_KEY_MESSAGE.to_string())); - } -} diff --git a/litellm-rust/crates/core/src/providers/mod.rs b/litellm-rust/crates/core/src/providers/mod.rs index c0c2c69831b..1aeb75063d6 100644 --- a/litellm-rust/crates/core/src/providers/mod.rs +++ b/litellm-rust/crates/core/src/providers/mod.rs @@ -2,7 +2,4 @@ pub mod anthropic; pub mod azure_ai; #[cfg(feature = "bedrock-auth")] pub mod bedrock; -pub mod mistral; pub mod openai; -pub mod reducto; -pub mod vertex_ai; diff --git a/litellm-rust/crates/core/src/providers/reducto/mod.rs b/litellm-rust/crates/core/src/providers/reducto/mod.rs deleted file mode 100644 index 3621ff6a2fd..00000000000 --- a/litellm-rust/crates/core/src/providers/reducto/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod ocr; diff --git a/litellm-rust/crates/core/src/providers/reducto/ocr/mod.rs b/litellm-rust/crates/core/src/providers/reducto/ocr/mod.rs deleted file mode 100644 index 8acee8f770c..00000000000 --- a/litellm-rust/crates/core/src/providers/reducto/ocr/mod.rs +++ /dev/null @@ -1,4 +0,0 @@ -pub mod transformation; - -#[cfg(test)] -mod tests; diff --git a/litellm-rust/crates/core/src/providers/reducto/ocr/tests.rs b/litellm-rust/crates/core/src/providers/reducto/ocr/tests.rs deleted file mode 100644 index 2b66d058b5d..00000000000 --- a/litellm-rust/crates/core/src/providers/reducto/ocr/tests.rs +++ /dev/null @@ -1,202 +0,0 @@ -use rstest::{fixture, rstest}; -use serde_json::{Value, json}; - -use super::transformation::*; -use crate::ocr::transformation::OcrProviderConfig; - -#[fixture] -fn parse_response() -> Value { - json!({ - "job_id": "job_123", - "usage": {"num_pages": 3, "credits": 3}, - "result": { - "chunks": [ - { - "content": "Page 1 block A", - "blocks": [{ - "content": "Page 1 block A", - "bbox": {"page": 1}, - "kind": "text", - }], - }, - { - "content": "Page 2 block A", - "blocks": [{ - "content": "Page 2 block A", - "bbox": {"page": 2}, - "kind": "table", - }], - }, - { - "content": "Page 1 block B", - "blocks": [{ - "content": "Page 1 block B", - "bbox": {"page": 1}, - "kind": "text", - }], - }, - { - "content": "Page 3 block A", - "blocks": [{ - "content": "Page 3 block A", - "bbox": {"page": 3}, - "kind": "figure", - }], - }, - ], - }, - }) -} - -#[rstest] -fn test_parse_v3_file_upload_and_response_mapping(parse_response: Value) { - let source = classify_document_source("data:application/pdf;base64,JVBERi0xLjQ=") - .expect("PDF data URI should be valid"); - let upload = build_upload_request( - source, - "Bearer test-key", - Some("https://platform.reducto.ai"), - ) - .expect("data URI should require upload"); - assert_eq!(upload.url, "https://platform.reducto.ai/upload"); - assert_eq!(upload.authorization, "Bearer test-key"); - assert_eq!(upload.file_name, "document"); - assert_eq!(upload.mime_type, "application/pdf"); - assert_eq!(upload.bytes, b"%PDF-1.4"); - - let optional_params = json!({ - "formatting": {"table_output_format": "html"}, - "retrieval": {"chunk_mode": "section"}, - "settings": {"ocr_system": "standard"}, - }) - .as_object() - .expect("params should be an object") - .clone(); - let request = build_parse_v3_request("reducto://uploaded.pdf", optional_params); - assert_eq!( - request.data, - json!({ - "input": "reducto://uploaded.pdf", - "formatting": {"table_output_format": "html"}, - "retrieval": {"chunk_mode": "section"}, - "settings": {"ocr_system": "standard"}, - }) - ); - - let transformed = transform_reducto_response("parse-v3", parse_response.clone()) - .expect("response should transform"); - assert_eq!( - transformed.usage_info, - Some(json!({"pages_processed": 3, "credits": 3})) - ); - assert_eq!(transformed.pages.len(), 3); - assert_eq!( - transformed.pages[0], - json!({ - "index": 0, - "markdown": "Page 1 block A\n\nPage 1 block B", - "blocks": [ - {"content": "Page 1 block A", "bbox": {"page": 1}, "kind": "text"}, - {"content": "Page 1 block B", "bbox": {"page": 1}, "kind": "text"}, - ], - }) - ); - assert_eq!(transformed.pages[1]["markdown"], "Page 2 block A"); - assert_eq!(transformed.pages[2]["markdown"], "Page 3 block A"); - assert_eq!(transformed.provider_native_response, Some(parse_response)); -} - -#[rstest] -fn test_parse_v3_reducto_id_passthrough_skips_upload(parse_response: Value) { - let document = json!({ - "type": "document_url", - "document_url": "reducto://already-uploaded.pdf", - }); - let source = extract_document_source(&document).expect("Reducto ID should be valid"); - assert!(build_upload_request(source.clone(), "Bearer test-key", None).is_none()); - assert_eq!( - source, - ReductoDocumentSource::FileId("reducto://already-uploaded.pdf".to_string()) - ); - - let request = REDUCTO_PARSE_V3_CONFIG - .transform_ocr_request( - "parse-v3", - document, - json!({"retrieval": {"chunk_mode": "section"}}) - .as_object() - .expect("params should be object") - .clone(), - ) - .expect("direct ID should transform"); - assert_eq!(request.data["input"], "reducto://already-uploaded.pdf"); - assert_eq!(request.data["retrieval"]["chunk_mode"], "section"); - - let response = REDUCTO_PARSE_V3_CONFIG - .transform_ocr_response("parse-v3", parse_response) - .expect("response should transform"); - assert!( - response.pages[0]["markdown"] - .as_str() - .expect("markdown should be string") - .starts_with("Page 1 block A") - ); -} - -#[rstest] -fn test_parse_legacy_wraps_enhance_under_options() { - let request = build_parse_legacy_request( - "reducto://legacy.pdf", - json!({"enhance": {"agentic": [{"type": "table"}]}}) - .as_object() - .expect("params should be object"), - ); - assert_eq!( - request.data, - json!({ - "document_url": "reducto://legacy.pdf", - "options": {"enhance": {"agentic": [{"type": "table"}]}}, - }) - ); -} - -#[rstest] -fn test_parse_v3_image_data_uri_upload_uses_image_mime() { - let source = classify_document_source("data:image/png;base64,iVBORw0KGgo=") - .expect("PNG data URI should be valid"); - let upload = build_upload_request( - source, - "Bearer programmatic-key", - Some("https://custom.reducto.test/"), - ) - .expect("data URI should require upload"); - assert_eq!(upload.url, "https://custom.reducto.test/upload"); - assert_eq!(upload.authorization, "Bearer programmatic-key"); - assert_eq!(upload.mime_type, "image/png"); - assert_eq!(upload.bytes, b"\x89PNG\r\n\x1a\n"); -} - -#[rstest] -#[case::http("http://example.com/document.pdf")] -#[case::https("https://example.com/document.pdf")] -fn test_parse_v3_rejects_plain_http_urls(#[case] source: &str) { - let error = classify_document_source(source).expect_err("plain URL should be rejected"); - assert!(error.to_string().contains("upload the file first")); -} - -#[rstest] -fn test_parse_v3_uses_programmatic_api_key_over_env() { - let key = resolve_api_key(Some("passed-key"), &|_| Some("env-reducto-key".to_string())) - .expect("explicit key should resolve"); - assert_eq!(key, "passed-key"); - - let headers = REDUCTO_PARSE_V3_CONFIG - .validate_environment(Vec::new(), Some("passed-key"), &|_| { - Some("env-reducto-key".to_string()) - }) - .expect("headers should validate"); - assert_eq!( - headers, - vec![("Authorization".to_string(), "Bearer passed-key".to_string())] - ); -} diff --git a/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs deleted file mode 100644 index f025887846f..00000000000 --- a/litellm-rust/crates/core/src/providers/reducto/ocr/transformation.rs +++ /dev/null @@ -1,407 +0,0 @@ -use std::collections::BTreeMap; - -use base64::Engine; -use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; -use serde_json::{Map, Value, json}; - -use crate::error::{Error, json_type_name}; -use crate::ocr::transformation::OcrProviderConfig; -use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData}; - -pub const REDUCTO_API_BASE: &str = "https://platform.reducto.ai"; -pub const REDUCTO_API_KEY_ENV: &str = "REDUCTO_API_KEY"; -pub const REDUCTO_ID_PREFIX: &str = "reducto://"; - -const PARSE_V3_SUPPORTED_OCR_PARAMS: &[&str] = &["formatting", "retrieval", "settings"]; -const PARSE_LEGACY_SUPPORTED_OCR_PARAMS: &[&str] = &["enhance"]; -const MISSING_KEY_MESSAGE: &str = "Missing REDUCTO_API_KEY - set it in the environment or pass api_key to litellm.ocr()/litellm.aocr()"; -const DATA_URI_UPLOAD_REQUIRED: &str = - "Reducto data URI upload must complete before OCR request transformation"; - -#[derive(Clone, Debug, PartialEq, Eq)] -pub enum ReductoDocumentSource { - FileId(String), - Upload { bytes: Vec, mime_type: String }, -} - -#[derive(Clone, PartialEq, Eq)] -pub struct ReductoUploadRequest { - pub url: String, - pub authorization: String, - pub file_name: &'static str, - pub bytes: Vec, - pub mime_type: String, -} - -pub struct ReductoParseV3Config; -pub struct ReductoParseLegacyConfig; - -pub const REDUCTO_PARSE_V3_CONFIG: ReductoParseV3Config = ReductoParseV3Config; -pub const REDUCTO_PARSE_LEGACY_CONFIG: ReductoParseLegacyConfig = ReductoParseLegacyConfig; - -pub fn config_for_model(model: &str) -> Option<&'static dyn OcrProviderConfig> { - match model { - "parse-v3" => Some(&REDUCTO_PARSE_V3_CONFIG), - "parse-legacy" => Some(&REDUCTO_PARSE_LEGACY_CONFIG), - _ => None, - } -} - -pub fn normalize_api_base(api_base: Option<&str>) -> String { - api_base - .map(str::trim) - .filter(|base| !base.is_empty()) - .unwrap_or(REDUCTO_API_BASE) - .trim_end_matches('/') - .to_string() -} - -pub fn parse_url(api_base: Option<&str>) -> String { - format!("{}/parse", normalize_api_base(api_base)) -} - -pub fn upload_url(api_base: Option<&str>) -> String { - format!("{}/upload", normalize_api_base(api_base)) -} - -pub fn resolve_api_key( - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - api_key - .map(str::trim) - .filter(|key| !key.is_empty()) - .map(str::to_string) - .or_else(|| { - env_lookup(REDUCTO_API_KEY_ENV) - .map(|key| key.trim().to_string()) - .filter(|key| !key.is_empty()) - }) - .ok_or_else(|| Error::Auth(MISSING_KEY_MESSAGE.to_string())) -} - -pub fn extract_document_source(document: &Value) -> Result { - let document = document.as_object().ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(document), - })?; - let source = document - .get("document_url") - .and_then(Value::as_str) - .filter(|source| !source.is_empty()) - .or_else(|| document.get("image_url").and_then(Value::as_str)) - .ok_or_else(|| { - Error::InvalidRequest( - "Reducto expected OCR preprocessing to produce document_url or image_url" - .to_string(), - ) - })?; - classify_document_source(source) -} - -pub fn classify_document_source(source: &str) -> Result { - if source.starts_with(REDUCTO_ID_PREFIX) { - return Ok(ReductoDocumentSource::FileId(source.to_string())); - } - if source.starts_with("http://") || source.starts_with("https://") { - return Err(Error::InvalidRequest( - "Reducto requires type='file' (auto-uploaded) or a reducto:// id. Plain http(s) URLs are not supported; upload the file first." - .to_string(), - )); - } - if !source.starts_with("data:") { - return Err(Error::InvalidRequest( - "Reducto requires a reducto:// id or a base64 data URI after OCR preprocessing." - .to_string(), - )); - } - - let (header, encoded) = source - .split_once(',') - .ok_or_else(|| Error::InvalidRequest("Invalid Reducto data URI provided.".to_string()))?; - if !header.split(';').any(|part| part == "base64") { - return Err(Error::InvalidRequest( - "Reducto only supports base64-encoded data URIs.".to_string(), - )); - } - - let mime_type = header - .strip_prefix("data:") - .and_then(|header| header.split(';').next()) - .filter(|mime| !mime.is_empty()) - .unwrap_or("application/octet-stream") - .to_string(); - let bytes = BASE64_STANDARD.decode(encoded).map_err(|_| { - Error::InvalidRequest("Invalid Reducto base64 payload provided.".to_string()) - })?; - - Ok(ReductoDocumentSource::Upload { bytes, mime_type }) -} - -pub fn build_upload_request( - source: ReductoDocumentSource, - authorization: &str, - api_base: Option<&str>, -) -> Option { - let ReductoDocumentSource::Upload { bytes, mime_type } = source else { - return None; - }; - - Some(ReductoUploadRequest { - url: upload_url(api_base), - authorization: authorization.to_string(), - file_name: "document", - bytes, - mime_type, - }) -} - -pub fn extract_upload_file_id(response_json: &Value) -> Result<&str, Error> { - response_json - .as_object() - .and_then(|response| response.get("file_id")) - .and_then(Value::as_str) - .filter(|file_id| !file_id.is_empty()) - .ok_or_else(|| { - Error::InvalidResponse(format!( - "Reducto /upload returned 200 without a file_id; got payload={response_json}" - )) - }) -} - -pub fn build_parse_v3_request( - file_id: &str, - optional_params: Map, -) -> OcrRequestData { - let data = std::iter::once(("input".to_string(), Value::String(file_id.to_string()))) - .chain(optional_params) - .collect(); - OcrRequestData { - data: Value::Object(data), - files: None, - } -} - -pub fn build_parse_legacy_request( - file_id: &str, - optional_params: &Map, -) -> OcrRequestData { - let options = optional_params - .get("enhance") - .filter(|enhance| !enhance.is_null()) - .map(|enhance| json!({"options": {"enhance": enhance}})); - let data = match options { - Some(Value::Object(options)) => std::iter::once(( - "document_url".to_string(), - Value::String(file_id.to_string()), - )) - .chain(options) - .collect(), - _ => Map::from_iter([( - "document_url".to_string(), - Value::String(file_id.to_string()), - )]), - }; - OcrRequestData { - data: Value::Object(data), - files: None, - } -} - -fn source_file_id(document: &Value) -> Result { - match extract_document_source(document)? { - ReductoDocumentSource::FileId(file_id) => Ok(file_id), - ReductoDocumentSource::Upload { .. } => Err(Error::Unsupported(DATA_URI_UPLOAD_REQUIRED)), - } -} - -fn page_number(block: &Map) -> Option { - let page = block.get("bbox")?.as_object()?.get("page")?; - page.as_i64() - .or_else(|| page.as_u64().and_then(|page| i64::try_from(page).ok())) - .or_else(|| page.as_str().and_then(|page| page.parse().ok())) -} - -fn chunks(result: &Map) -> &[Value] { - result - .get("chunks") - .and_then(Value::as_array) - .map(Vec::as_slice) - .unwrap_or_default() -} - -fn build_pages(result: &Map) -> Vec { - let blocks_by_page = chunks(result) - .iter() - .filter_map(Value::as_object) - .filter_map(|chunk| chunk.get("blocks").and_then(Value::as_array)) - .flatten() - .filter_map(|block| block.as_object().map(|object| (block, object))) - .filter_map(|(block, object)| page_number(object).map(|page| (page, block.clone()))) - .fold( - BTreeMap::>::new(), - |mut pages, (page, block)| { - pages.entry(page).or_default().push(block); - pages - }, - ); - - if blocks_by_page.is_empty() { - let markdown = chunks(result) - .iter() - .filter_map(Value::as_object) - .filter_map(|chunk| chunk.get("content").and_then(Value::as_str)) - .filter(|content| !content.is_empty()) - .collect::>() - .join("\n\n"); - return if markdown.is_empty() { - Vec::new() - } else { - vec![json!({"index": 0, "markdown": markdown})] - }; - } - - blocks_by_page - .into_iter() - .map(|(page, blocks)| { - let markdown = blocks - .iter() - .filter_map(Value::as_object) - .filter_map(|block| block.get("content").and_then(Value::as_str)) - .filter(|content| !content.is_empty()) - .collect::>() - .join("\n\n"); - json!({ - "index": page.saturating_sub(1).max(0), - "markdown": markdown, - "blocks": blocks, - }) - }) - .collect() -} - -pub fn transform_reducto_response( - model: &str, - response_json: Value, -) -> Result { - let response = response_json - .as_object() - .ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(&response_json), - })?; - let empty_result = Map::new(); - let result = match response.get("result") { - Some(Value::Object(result)) => result, - Some(Value::Null) => &empty_result, - Some(_) => { - return Err(Error::InvalidResponse( - "Reducto result must be an object".to_string(), - )); - } - None => response, - }; - let usage = response - .get("usage") - .and_then(Value::as_object) - .cloned() - .unwrap_or_default(); - let usage_info = Some(json!({ - "pages_processed": usage.get("num_pages").cloned().unwrap_or(Value::Null), - "credits": usage.get("credits").cloned().unwrap_or(Value::Null), - })); - - Ok(LiteLLMOcrResponse { - pages: build_pages(result), - model: model.to_string(), - document_annotation: None, - usage_info, - object: "ocr".to_string(), - extra_fields: Map::new(), - provider_native_response: Some(response_json), - }) -} - -impl OcrProviderConfig for ReductoParseV3Config { - fn supported_ocr_params(&self) -> &'static [&'static str] { - PARSE_V3_SUPPORTED_OCR_PARAMS - } - - fn transform_ocr_request( - &self, - _model: &str, - document: Value, - optional_params: Map, - ) -> Result { - let file_id = source_file_id(&document)?; - Ok(build_parse_v3_request(&file_id, optional_params)) - } - - fn transform_ocr_response( - &self, - model: &str, - response_json: Value, - ) -> Result { - transform_reducto_response(model, response_json) - } - - fn complete_url( - &self, - api_base: Option<&str>, - _model: &str, - _optional_params: &Map, - _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - Ok(parse_url(api_base)) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_api_key(api_key, env_lookup) - } -} - -impl OcrProviderConfig for ReductoParseLegacyConfig { - fn supported_ocr_params(&self) -> &'static [&'static str] { - PARSE_LEGACY_SUPPORTED_OCR_PARAMS - } - - fn transform_ocr_request( - &self, - _model: &str, - document: Value, - optional_params: Map, - ) -> Result { - let file_id = source_file_id(&document)?; - Ok(build_parse_legacy_request(&file_id, &optional_params)) - } - - fn transform_ocr_response( - &self, - model: &str, - response_json: Value, - ) -> Result { - transform_reducto_response(model, response_json) - } - - fn complete_url( - &self, - api_base: Option<&str>, - _model: &str, - _optional_params: &Map, - _env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - Ok(parse_url(api_base)) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_api_key(api_key, env_lookup) - } -} diff --git a/litellm-rust/crates/core/src/providers/vertex_ai/mod.rs b/litellm-rust/crates/core/src/providers/vertex_ai/mod.rs deleted file mode 100644 index 3621ff6a2fd..00000000000 --- a/litellm-rust/crates/core/src/providers/vertex_ai/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod ocr; diff --git a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/mod.rs b/litellm-rust/crates/core/src/providers/vertex_ai/ocr/mod.rs deleted file mode 100644 index f239b6921fa..00000000000 --- a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod transformation; diff --git a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs deleted file mode 100644 index b0e5a278d0b..00000000000 --- a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs +++ /dev/null @@ -1,467 +0,0 @@ -use crate::error::{Error, json_type_name}; -use crate::ocr::transformation::OcrProviderConfig; -use crate::ocr::types::{LiteLLMOcrResponse, OcrRequestData}; -use serde_json::{Map, Value, json}; - -use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; - -const VERTEX_DEFAULT_LOCATION: &str = "us-central1"; -const VERTEX_DEFAULT_DEEPSEEK_API_BASE: &str = "https://aiplatform.googleapis.com"; -const VERTEX_AI_API_KEY_ENV: &str = "VERTEX_AI_API_KEY"; -const VERTEXAI_API_KEY_ENV: &str = "VERTEXAI_API_KEY"; -const VERTEXAI_PROJECT_ENV: &str = "VERTEXAI_PROJECT"; -const VERTEXAI_LOCATION_ENV: &str = "VERTEXAI_LOCATION"; -const VERTEX_LOCATION_ENV: &str = "VERTEX_LOCATION"; - -#[rustfmt::skip] -const DEEPSEEK_SUPPORTED_OCR_PARAMS: &[&str] = &[ - "stream", - "temperature", - "max_tokens", - "top_p", - "n", - "stop", -]; - -pub struct VertexAiOcrConfig; -pub struct VertexAiDeepSeekOcrConfig; - -pub const VERTEX_AI_OCR_CONFIG: VertexAiOcrConfig = VertexAiOcrConfig; -pub const VERTEX_AI_DEEPSEEK_OCR_CONFIG: VertexAiDeepSeekOcrConfig = VertexAiDeepSeekOcrConfig; - -fn string_param<'a>(params: &'a Map, keys: &[&str]) -> Option<&'a str> { - keys.iter() - .find_map(|key| params.get(*key).and_then(Value::as_str)) - .map(str::trim) - .filter(|value| !value.is_empty()) -} - -pub fn is_deepseek_model(model: &str) -> bool { - model.to_ascii_lowercase().contains("deepseek") -} - -pub fn resolve_vertex_api_key( - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - api_key - .map(str::trim) - .filter(|key| !key.is_empty()) - .map(str::to_string) - .or_else(|| env_lookup(VERTEX_AI_API_KEY_ENV).filter(|key| !key.trim().is_empty())) - .or_else(|| env_lookup(VERTEXAI_API_KEY_ENV).filter(|key| !key.trim().is_empty())) - .ok_or_else(|| { - Error::Auth( - "Missing Vertex AI access token - pass api_key or provide Authorization via extra_headers" - .to_string(), - ) - }) -} - -fn vertex_project( - params: &Map, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - string_param(params, &["vertex_project", "vertex_ai_project"]) - .map(str::to_string) - .or_else(|| env_lookup(VERTEXAI_PROJECT_ENV).filter(|value| !value.trim().is_empty())) - .ok_or_else(|| { - Error::InvalidRequest( - "Missing vertex_project - Set VERTEXAI_PROJECT environment variable or pass vertex_project parameter" - .to_string(), - ) - }) -} - -fn vertex_location( - params: &Map, - env_lookup: &dyn Fn(&str) -> Option, -) -> String { - string_param(params, &["vertex_location", "vertex_ai_location"]) - .map(str::to_string) - .or_else(|| env_lookup(VERTEXAI_LOCATION_ENV).filter(|value| !value.trim().is_empty())) - .or_else(|| env_lookup(VERTEX_LOCATION_ENV).filter(|value| !value.trim().is_empty())) - .unwrap_or_else(|| VERTEX_DEFAULT_LOCATION.to_string()) -} - -fn vertex_mistral_api_base(api_base: Option<&str>, location: &str) -> String { - api_base - .map(str::trim) - .filter(|value| !value.is_empty()) - .map(str::to_string) - .unwrap_or_else(|| format!("https://{location}-aiplatform.googleapis.com")) - .trim_end_matches('/') - .to_string() -} - -pub fn complete_vertex_mistral_url( - api_base: Option<&str>, - model: &str, - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - let project = vertex_project(optional_params, env_lookup)?; - let location = vertex_location(optional_params, env_lookup); - let base = vertex_mistral_api_base(api_base, &location); - Ok(format!( - "{base}/v1/projects/{project}/locations/{location}/publishers/mistralai/models/{model}:rawPredict" - )) -} - -pub fn complete_vertex_deepseek_url( - api_base: Option<&str>, - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, -) -> Result { - let project = vertex_project(optional_params, env_lookup)?; - let location = vertex_location(optional_params, env_lookup); - let base = api_base - .map(str::trim) - .filter(|value| !value.is_empty()) - .unwrap_or(VERTEX_DEFAULT_DEEPSEEK_API_BASE) - .trim_end_matches('/'); - Ok(format!( - "{base}/v1/projects/{project}/locations/{location}/endpoints/openapi/chat/completions" - )) -} - -fn document_content_item(document: &Value) -> Result { - let object = document.as_object().ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(document), - })?; - let doc_type = object - .get("type") - .and_then(Value::as_str) - .ok_or(Error::MissingField("document.type"))?; - let url_field = match doc_type { - "image_url" => "image_url", - "document_url" => "document_url", - other => { - return Err(Error::InvalidRequest(format!( - "Unsupported document type: {other}. Expected 'image_url' or 'document_url'" - ))); - } - }; - let url = object - .get(url_field) - .and_then(Value::as_str) - .filter(|value| !value.is_empty()) - .ok_or(Error::MissingField(url_field))?; - - Ok(json!({ - "type": "image_url", - "image_url": url, - })) -} - -fn deepseek_model_name(model: &str) -> String { - if model.starts_with("deepseek-ai/") { - model.to_string() - } else { - format!("deepseek-ai/{model}") - } -} - -fn first_choice_content(response: &Value) -> Result { - response - .get("choices") - .and_then(Value::as_array) - .and_then(|choices| choices.first()) - .and_then(|choice| choice.get("message")) - .and_then(|message| message.get("content")) - .cloned() - .filter(|content| match content { - Value::String(value) => !value.is_empty(), - Value::Object(_) => true, - _ => false, - }) - .ok_or_else(|| Error::InvalidResponse("No content in DeepSeek OCR response".to_string())) -} - -fn ocr_data_from_content(content: Value, usage: Option, model: &str) -> Value { - match content { - Value::String(content) => { - if content.trim_start().starts_with('{') { - serde_json::from_str(&content).unwrap_or_else(|_| { - json!({ - "pages": [{"index": 0, "markdown": content}], - "model": model, - "usage_info": usage.unwrap_or_else(|| json!({})), - }) - }) - } else { - json!({ - "pages": [{"index": 0, "markdown": content}], - "model": model, - "usage_info": usage.unwrap_or_else(|| json!({})), - }) - } - } - Value::Object(_) => content, - other => json!({ - "pages": [{"index": 0, "markdown": other.to_string()}], - "model": model, - "usage_info": usage.unwrap_or_else(|| json!({})), - }), - } -} - -impl OcrProviderConfig for VertexAiOcrConfig { - fn supported_ocr_params(&self) -> &'static [&'static str] { - MISTRAL_OCR_CONFIG.supported_ocr_params() - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn transform_ocr_request( - &self, - model: &str, - document: Value, - optional_params: Map, - ) -> Result { - MISTRAL_OCR_CONFIG.transform_ocr_request(model, document, optional_params) - } - - fn transform_ocr_response( - &self, - model: &str, - response_json: Value, - ) -> Result { - MISTRAL_OCR_CONFIG.transform_ocr_response(model, response_json) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn complete_url( - &self, - api_base: Option<&str>, - model: &str, - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - complete_vertex_mistral_url(api_base, model, optional_params, env_lookup) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_vertex_api_key(api_key, env_lookup) - } - - fn requires_data_uri_document(&self) -> bool { - true - } -} - -impl OcrProviderConfig for VertexAiDeepSeekOcrConfig { - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn supported_ocr_params(&self) -> &'static [&'static str] { - DEEPSEEK_SUPPORTED_OCR_PARAMS - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn map_ocr_params(&self, non_default_params: &Map) -> Map { - non_default_params - .iter() - .filter(|(name, _)| DEEPSEEK_SUPPORTED_OCR_PARAMS.contains(&name.as_str())) - .map(|(name, value)| (name.clone(), value.clone())) - .collect() - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn transform_ocr_request( - &self, - model: &str, - document: Value, - optional_params: Map, - ) -> Result { - let mut data = Map::new(); - data.insert( - "model".to_string(), - Value::String(deepseek_model_name(model)), - ); - data.insert( - "messages".to_string(), - json!([{"role": "user", "content": [document_content_item(&document)?]}]), - ); - for (key, value) in optional_params { - if DEEPSEEK_SUPPORTED_OCR_PARAMS.contains(&key.as_str()) { - data.insert(key, value); - } - } - Ok(OcrRequestData { - data: Value::Object(data), - files: None, - }) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn transform_ocr_response( - &self, - model: &str, - response_json: Value, - ) -> Result { - let response = response_json - .as_object() - .ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(&response_json), - })?; - let usage = response.get("usage").cloned(); - let content = first_choice_content(&response_json)?; - let mut ocr_data = ocr_data_from_content(content.clone(), usage.clone(), model); - - if !ocr_data.get("pages").is_some_and(Value::is_array) { - ocr_data = json!({ - "pages": [{ - "index": 0, - "markdown": match content { - Value::String(value) => value, - other => other.to_string(), - } - }], - "model": ocr_data.get("model").and_then(Value::as_str).unwrap_or(model), - "usage_info": ocr_data.get("usage_info").cloned().or(usage).unwrap_or_else(|| json!({})), - }); - } - - let object = ocr_data.as_object().ok_or_else(|| Error::InvalidType { - expected: "object", - actual: json_type_name(&ocr_data), - })?; - let pages = object - .get("pages") - .and_then(Value::as_array) - .cloned() - .unwrap_or_default(); - let usage_info = object - .get("usage_info") - .cloned() - .or_else(|| response.get("usage").cloned()); - Ok(LiteLLMOcrResponse { - pages, - model: object - .get("model") - .and_then(Value::as_str) - .unwrap_or(model) - .to_string(), - document_annotation: object.get("document_annotation").cloned(), - usage_info, - object: "ocr".to_string(), - extra_fields: Map::new(), - provider_native_response: None, - }) - } - - #[tracing::instrument(target = "litellm::function_trace", level = "trace", skip_all)] - fn complete_url( - &self, - api_base: Option<&str>, - _model: &str, - optional_params: &Map, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - complete_vertex_deepseek_url(api_base, optional_params, env_lookup) - } - - fn resolve_api_key( - &self, - api_key: Option<&str>, - env_lookup: &dyn Fn(&str) -> Option, - ) -> Result { - resolve_vertex_api_key(api_key, env_lookup) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use rstest::rstest; - - #[test] - fn vertex_mistral_url_uses_project_location_and_model() { - let params = Map::from_iter([ - ("vertex_project".to_string(), json!("proj-1")), - ("vertex_location".to_string(), json!("europe-west4")), - ]); - - let url = complete_vertex_mistral_url(None, "mistral-ocr-maas", ¶ms, &|_| None) - .expect("url builds"); - - assert_eq!( - url, - "https://europe-west4-aiplatform.googleapis.com/v1/projects/proj-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict" - ); - } - - #[test] - fn vertex_mistral_reuses_mistral_body_transform() { - let body = VERTEX_AI_OCR_CONFIG - .transform_ocr_request( - "mistral-ocr-maas", - json!({"type": "image_url", "image_url": "data:image/png;base64,abc"}), - Map::new(), - ) - .expect("request transforms") - .data; - - assert_eq!(body["model"], "mistral-ocr-maas"); - assert_eq!(body["document"]["image_url"], "data:image/png;base64,abc"); - } - - #[test] - fn vertex_deepseek_request_uses_ocr_endpoint_shape() { - let body = VERTEX_AI_DEEPSEEK_OCR_CONFIG - .transform_ocr_request( - "deepseek-ocr-maas", - json!({"type": "document_url", "document_url": "gs://bucket/doc.pdf"}), - Map::from_iter([("temperature".to_string(), json!(0.1))]), - ) - .expect("request transforms") - .data; - - assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas"); - assert_eq!(body["temperature"], 0.1); - assert_eq!( - body["messages"][0]["content"][0], - json!({"type": "image_url", "image_url": "gs://bucket/doc.pdf"}) - ); - } - - #[rstest] - #[case::bare_model("deepseek-ocr-maas")] - #[case::namespaced_model("deepseek-ai/deepseek-ocr-maas")] - fn vertex_deepseek_request_uses_single_provider_namespace(#[case] model: &str) { - let body = VERTEX_AI_DEEPSEEK_OCR_CONFIG - .transform_ocr_request( - model, - json!({"type": "image_url", "image_url": "data:image/png;base64,AA=="}), - Map::new(), - ) - .expect("request transforms") - .data; - - assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas"); - } - - #[test] - fn vertex_deepseek_response_wraps_markdown_content() { - let response = VERTEX_AI_DEEPSEEK_OCR_CONFIG - .transform_ocr_response( - "deepseek-ocr-maas", - json!({ - "choices": [{"message": {"content": "# OCR text"}}], - "usage": {"prompt_tokens": 1} - }), - ) - .expect("response transforms"); - - assert_eq!( - response.pages, - vec![json!({"index": 0, "markdown": "# OCR text"})] - ); - assert_eq!(response.model, "deepseek-ocr-maas"); - assert_eq!(response.usage_info, Some(json!({"prompt_tokens": 1}))); - } -} diff --git a/litellm-rust/crates/core/src/url_utils.rs b/litellm-rust/crates/core/src/url_utils.rs index 982dca0dbe3..1150f93a5c7 100644 --- a/litellm-rust/crates/core/src/url_utils.rs +++ b/litellm-rust/crates/core/src/url_utils.rs @@ -60,6 +60,14 @@ impl ApiUrl { } impl ApiUrl { + pub(crate) fn append_query_pairs<'a>( + mut self, + pairs: impl IntoIterator, + ) -> Self { + self.url.query_pairs_mut().extend_pairs(pairs); + self + } + pub(crate) fn into_string(self) -> String { self.url.into() } @@ -92,4 +100,19 @@ mod tests { .expect("url builds"); assert_eq!(actual, "https://example.test/v1/ocr?tenant=a"); } + + #[test] + fn appended_query_pairs_are_encoded() { + let actual = ApiUrl::parse("https://example.test") + .and_then(|url| url.complete_path(&["analyze"])) + .map(|url| { + url.append_query_pairs([("model", "name with spaces")]) + .into_string() + }) + .expect("url builds"); + assert_eq!( + actual, + "https://example.test/analyze?model=name+with+spaces" + ); + } } diff --git a/litellm-rust/crates/core/tests/azure_ai_ocr.rs b/litellm-rust/crates/core/tests/azure_ai_ocr.rs new file mode 100644 index 00000000000..d7d532cfef1 --- /dev/null +++ b/litellm-rust/crates/core/tests/azure_ai_ocr.rs @@ -0,0 +1,97 @@ +use std::sync::Arc; + +use serde_json::{Value, json}; + +use super::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks}; +use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; + +#[tokio::test] +async fn facade_executes_azure_mistral_with_prepared_auth() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ + "pages":[{"index":0,"markdown":"hello"}], + "usage_info":{"pages_processed":1} + }))]) + .await; + let mut request = wire_request( + "azure_ai/model", + &base, + json!({"include_image_base64":true}), + ); + request.connection.api_key = None; + request.connection.extra_headers = vec![( + "Authorization".into(), + "Bearer python-prepared-token".into(), + )]; + + let result = perform_ocr(request).await.unwrap(); + server.await.unwrap(); + assert_eq!(result.pages[0]["markdown"], "hello"); + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert!(requests[0].starts_with("POST /providers/mistral/azure/ocr ")); + assert!( + requests[0] + .to_ascii_lowercase() + .contains("authorization: bearer python-prepared-token\r\n") + ); + let body: Value = serde_json::from_str(requests[0].split_once("\r\n\r\n").unwrap().1).unwrap(); + assert_eq!( + body, + json!({ + "model":"model", + "document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}, + "include_image_base64":true + }) + ); +} + +#[tokio::test] +async fn facade_acquires_supplied_entra_token_for_final_request() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; + let mut request = wire_request( + "azure_ai/model", + &base, + json!({"azure_ad_token":"rust-owned-token"}), + ); + request.connection.api_key = None; + + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert!( + requests[0] + .to_ascii_lowercase() + .contains("authorization: bearer rust-owned-token\r\n") + ); +} + +struct ReplaceBodyDocument; + +impl OcrHooks for ReplaceBodyDocument { + fn has_guardrails(&self) -> bool { + true + } + + fn during_call( + &self, + mut request: OcrDuringCallRequest, + ) -> OcrHookFuture<'_, OcrDuringCallRequest> { + Box::pin(async move { + request.body["document"] = json!({ + "type":"document_url", + "document_url":"https://example.com/not-inline.pdf" + }); + Ok(request) + }) + } +} + +#[tokio::test] +async fn rejects_non_inline_body_after_guardrails() { + let mut request = wire_request("azure_ai/model", "http://127.0.0.1:1", json!({})); + request.hooks = Arc::new(ReplaceBodyDocument); + let error = perform_ocr(request).await.unwrap_err(); + assert!(error.to_string().contains("data URI")); +} diff --git a/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs new file mode 100644 index 00000000000..e4c81dea5a7 --- /dev/null +++ b/litellm-rust/crates/core/tests/azure_document_intelligence_ocr.rs @@ -0,0 +1,395 @@ +use serde_json::{Value, json}; + +use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; +use super::wire::{OcrWireRequest, decode_request}; + +fn query_value(url: &str, key: &str) -> Option { + url::Url::parse(url) + .unwrap() + .query_pairs() + .find_map(|(name, value)| (name == key).then(|| value.into_owned())) +} + +#[tokio::test] +async fn facade_maps_pages_features_and_url_document() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ + "status":"succeeded", + "analyzeResult":{"pages":[]} + }))]) + .await; + let mut request = wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + &base, + json!({"pages":[2,0,0,1],"features":["keyValuePairs","languages"]}), + ); + request.document = serde_json::from_value(json!({ + "type":"document_url", + "document_url":"https://example.com/document.pdf" + })) + .unwrap(); + + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + let request = &seen.lock().unwrap()[0]; + let target = request.split_whitespace().nth(1).unwrap(); + let url = format!("{base}{target}"); + assert_eq!(query_value(&url, "pages").as_deref(), Some("1,2,3")); + assert_eq!( + query_value(&url, "features").as_deref(), + Some("keyValuePairs,languages") + ); + let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap(); + assert_eq!( + body, + json!({"urlSource":"https://example.com/document.pdf"}) + ); +} + +#[tokio::test] +async fn rejects_invalid_pages_features_and_format() { + for options in [ + json!({"pages":[true]}), + json!({"pages":[1,"2"]}), + json!({"pages":[-1]}), + json!({"pages":"1&&features=bad"}), + json!({"features":"languages&pages=1"}), + json!({"req_format":"azure"}), + ] { + let result = decode_request(OcrWireRequest { + model: "azure_ai/doc-intelligence/prebuilt-read".into(), + document: json!({"type":"document_url","document_url":"https://example.com/a.pdf"}), + api_key: Some("key".into()), + api_base: Some("http://127.0.0.1:1".into()), + custom_llm_provider: None, + extra_headers: None, + optional_params: options.as_object().unwrap().clone(), + input_sources: Default::default(), + timeout_seconds: None, + }); + let rejected = match result { + Ok(request) => perform_ocr(request).await.is_err(), + Err(_) => true, + }; + assert!(rejected, "accepted {options}"); + } +} + +#[tokio::test] +async fn inline_document_decodes_to_base64_source() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ + "status":"succeeded" + }))]) + .await; + let request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})); + + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + let request = &seen.lock().unwrap()[0]; + let body: Value = serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap(); + assert_eq!(body, json!({"base64Source":"YWJj"})); +} + +#[tokio::test] +async fn immediate_response_normalizes_pages_and_preserves_native() { + let operation = json!({ + "status":"succeeded", + "operationExtension":42, + "analyzeResult":{ + "content":"A\n\nB", + "tables":[{"cells":[]}], + "keyValuePairs":[{"key":{"content":"A"}}], + "pages":[{ + "pageNumber":"2", + "width":"8.5", + "height":11, + "unit":"inch", + "lines":[{"content":"A"},{"content":null},{"content":"B"}] + }] + } + }); + let (base, _, server) = mock_server(vec![MockResponse::json(operation.clone())]).await; + let result = perform_ocr(wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + &base, + json!({"req_format":"native"}), + )) + .await + .unwrap(); + server.await.unwrap(); + + assert_eq!(result.pages[0]["index"], 1); + assert_eq!(result.pages[0]["markdown"], "A\n\nB"); + assert_eq!( + result.pages[0]["dimensions"], + json!({"width":816,"height":1056,"dpi":96}) + ); + assert_eq!(result.usage_info, Some(json!({"pages_processed":1}))); + assert_eq!(result.provider_native_response, Some(operation)); +} + +#[tokio::test] +async fn accepted_response_polls_to_success_with_only_credentials() { + let operation = json!({"status":"succeeded","analyzeResult":{"pages":[]}}); + let (base, seen, server) = mock_server(vec![ + MockResponse { + status: 202, + headers: vec![("Operation-Location", "{base}/operation".into())], + body: json!({}), + }, + MockResponse { + status: 200, + headers: vec![("Retry-After", "0".into())], + body: json!({"status":"running"}), + }, + MockResponse::json(operation.clone()), + ]) + .await; + let mut request = wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + &base, + json!({"req_format":"native"}), + ); + request + .connection + .extra_headers + .push(("X-Trace".into(), "initial-only".into())); + + let result = perform_ocr(request).await.unwrap(); + server.await.unwrap(); + assert_eq!(result.provider_native_response, Some(operation)); + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 3); + assert!(requests[0].to_ascii_lowercase().contains("x-trace:")); + for poll in &requests[1..] { + assert!(!poll.to_ascii_lowercase().contains("x-trace:")); + assert!( + poll.to_ascii_lowercase() + .contains("ocp-apim-subscription-key: test-key") + ); + } +} + +#[tokio::test] +async fn polling_forwards_bearer_credentials() { + let (base, seen, server) = mock_server(vec![ + MockResponse { + status: 202, + headers: vec![("Operation-Location", "{base}/operation".into())], + body: json!({}), + }, + MockResponse::json(json!({"status":"succeeded"})), + ]) + .await; + let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})); + request.connection.api_key = None; + request.connection.extra_headers = vec![("Authorization".into(), "Bearer token".into())]; + + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + let requests = seen.lock().unwrap(); + assert!( + requests[1] + .to_ascii_lowercase() + .contains("authorization: bearer token") + ); +} + +#[tokio::test] +async fn polling_does_not_follow_redirects() { + let (base, seen, server) = mock_server(vec![ + MockResponse { + status: 202, + headers: vec![("Operation-Location", "{base}/operation".into())], + body: json!({}), + }, + MockResponse { + status: 302, + headers: vec![("Location", "{base}/redirected".into())], + body: json!({}), + }, + MockResponse::json(json!({"status":"succeeded"})), + ]) + .await; + + let error = perform_ocr(wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + &base, + json!({}), + )) + .await + .unwrap_err(); + + assert!(error.to_string().contains("status 302"), "{error}"); + assert_eq!(seen.lock().unwrap().len(), 2); + server.abort(); +} + +#[tokio::test] +async fn polling_rejects_terminal_failure() { + let (base, _, server) = mock_server(vec![ + MockResponse { + status: 202, + headers: vec![("Operation-Location", "{base}/operation".into())], + body: json!({}), + }, + MockResponse::json(json!({"status":"failed"})), + ]) + .await; + + let error = perform_ocr(wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + &base, + json!({}), + )) + .await + .unwrap_err(); + server.await.unwrap(); + assert!(error.to_string().contains("status failed")); +} + +#[tokio::test] +async fn malformed_provider_pages_report_response_paths() { + for (analysis, path) in [ + (json!({"pages":null}), "pages"), + (json!({"pages":[null]}), "pages[0]"), + (json!({"pages":[{"lines":null}]}), "lines"), + (json!({"pages":[{"width":"bad"}]}), "width"), + ] { + let (base, _, server) = mock_server(vec![MockResponse::json(json!({ + "status":"succeeded", + "analyzeResult":analysis + }))]) + .await; + let error = perform_ocr(wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + &base, + json!({}), + )) + .await + .unwrap_err(); + server.await.unwrap(); + assert!(error.to_string().contains(path), "{error}"); + } +} + +#[tokio::test] +async fn rejects_missing_invalid_and_cross_origin_operation_locations() { + for headers in [ + Vec::new(), + vec![("Operation-Location", "/relative".into())], + vec![("Operation-Location", "http://example.com/operation".into())], + vec![( + "Operation-Location", + "http://user:password@127.0.0.1/operation".into(), + )], + ] { + let (base, _, server) = mock_server(vec![MockResponse { + status: 202, + headers, + body: json!({}), + }]) + .await; + let error = perform_ocr(wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + &base, + json!({}), + )) + .await + .unwrap_err(); + server.await.unwrap(); + assert!(error.to_string().contains("operation-location")); + } +} + +#[tokio::test] +async fn polling_deadline_bounds_retry_delay() { + let (base, _, server) = mock_server(vec![ + MockResponse { + status: 202, + headers: vec![("Operation-Location", "{base}/operation".into())], + body: json!({}), + }, + MockResponse { + status: 200, + headers: vec![("Retry-After", "9999".into())], + body: json!({"status":"notStarted"}), + }, + ]) + .await; + let mut request = wire_request("azure_ai/doc-intelligence/prebuilt-read", &base, json!({})); + request.connection.poll_timeout = std::time::Duration::from_millis(100); + + let error = tokio::time::timeout(std::time::Duration::from_secs(1), perform_ocr(request)) + .await + .unwrap() + .unwrap_err(); + server.await.unwrap(); + assert!(error.to_string().contains("timed out")); +} + +#[tokio::test] +async fn model_id_is_encoded_and_dot_segments_are_rejected() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ + "status":"succeeded" + }))]) + .await; + perform_ocr(wire_request( + "azure_ai/doc-intelligence/a ?#é", + &base, + json!({}), + )) + .await + .unwrap(); + server.await.unwrap(); + assert!(seen.lock().unwrap()[0].contains("a%20%3F%23%C3%A9:analyze")); + + for model in [ + "azure_ai/doc-intelligence/.", + "azure_ai/doc-intelligence/..", + ] { + let error = perform_ocr(wire_request(model, "http://127.0.0.1:1", json!({}))) + .await + .unwrap_err(); + assert!(error.to_string().contains("dot segment")); + } +} + +#[tokio::test] +async fn pre_call_guardrail_receives_caller_pages_before_mapping() { + use crate::ocr::hooks::{OcrHookFuture, OcrHooks, OcrPreCallRequest}; + use std::sync::Arc; + + struct RewritePages; + impl OcrHooks for RewritePages { + fn has_guardrails(&self) -> bool { + true + } + + fn pre_call(&self, request: OcrPreCallRequest) -> OcrHookFuture<'_, OcrPreCallRequest> { + Box::pin(async move { + assert_eq!(request.optional_params["pages"], json!([0, 2])); + Ok(OcrPreCallRequest { + optional_params: json!({"pages": [1]}), + ..request + }) + }) + } + } + let (base, seen, server) = + mock_server(vec![MockResponse::json(json!({"status": "succeeded"}))]).await; + let request = wire_request( + "azure_ai/doc-intelligence/prebuilt-read", + &base, + json!({"pages": [0, 2]}), + ) + .with_host_hooks(Arc::new(RewritePages), None); + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + let requests = seen.lock().unwrap(); + let target = requests[0].split_whitespace().nth(1).unwrap(); + assert_eq!( + query_value(&format!("{base}{target}"), "pages").as_deref(), + Some("2") + ); + assert_eq!(requests.len(), 1); +} diff --git a/litellm-rust/crates/core/tests/deepseek_ocr.rs b/litellm-rust/crates/core/tests/deepseek_ocr.rs new file mode 100644 index 00000000000..875fc9e3dc6 --- /dev/null +++ b/litellm-rust/crates/core/tests/deepseek_ocr.rs @@ -0,0 +1,95 @@ +use rstest::rstest; +use serde_json::{Value, json}; + +use crate::ocr::codecs::deepseek::{ + DeepSeekOcrParams, DeepSeekOcrResponse, transform_ocr_request, transform_ocr_response, +}; +use crate::ocr::types::OcrDocument; + +fn document() -> OcrDocument { + serde_json::from_value(json!({"type":"image_url","image_url":"gs://bucket/a.png"})).unwrap() +} + +#[rstest] +#[case("stream", json!(true))] +#[case("temperature", json!(0.1))] +#[case("max_tokens", json!(1024))] +#[case("top_p", json!(0.9))] +#[case("n", json!(2))] +#[case("stop", json!("done"))] +#[case("stop", json!(["done", "stop"]))] +fn request_mapping_matches_python(#[case] name: &str, #[case] value: Value) { + let params: DeepSeekOcrParams = + serde_json::from_value(json!({name: value.clone(), "ignored": true})).unwrap(); + let result = serde_json::to_value( + transform_ocr_request("deepseek-ai/deepseek-ocr-maas", document(), ¶ms).unwrap(), + ) + .unwrap(); + assert_eq!(result["model"], "deepseek-ai/deepseek-ocr-maas"); + assert_eq!( + result["messages"][0]["content"][0], + json!({"type":"image_url","image_url":"gs://bucket/a.png"}) + ); + assert_eq!(result[name], value); + assert!(result.get("ignored").is_none()); +} + +#[rstest] +#[case(json!("# hello"), "# hello")] +#[case(json!("{broken"), "{broken")] +#[case(json!(" {\"pages\":[]} "), " {\"pages\":[]} ")] +#[case(json!({"pages":[]}), "{\"pages\":[]}")] +#[case(json!({}), "{}")] +#[case(json!("[]"), "[]")] +#[case(json!("{\"pages\":[{\"markdown\":\"json text\"}]}"), "json text")] +#[case(json!({"pages":[{"markdown":"object"}]}), "object")] +fn response_codec_handles_text_json_and_objects(#[case] content: Value, #[case] expected: &str) { + let response: DeepSeekOcrResponse = serde_json::from_value( + json!({"choices":[{"message":{"content":content}}],"usage":{"prompt_tokens":1}}), + ) + .unwrap(); + let result = transform_ocr_response("model", response) + .unwrap() + .into_json(); + assert_eq!(result["pages"][0]["markdown"], expected); + assert_eq!(result["pages"][0]["index"], 0); + assert_eq!(result["usage_info"]["prompt_tokens"], 1); +} + +#[test] +fn structured_result_maps_pages_usage_model_and_annotation() { + let response: DeepSeekOcrResponse = serde_json::from_value(json!({ + "choices":[{"message":{"content":{ + "pages":[{"index":2,"markdown":"page","images":[{"id":"one"}],"dimensions":{"width":10}}], + "model":"provider-model", + "usage_info":{"pages_processed":1}, + "document_annotation":{"language":"en"}, + "future":"kept" + }}}] + })) + .unwrap(); + let result = transform_ocr_response("requested", response) + .unwrap() + .into_json(); + assert_eq!(result["pages"][0]["index"], 2); + assert_eq!(result["pages"][0]["images"][0]["id"], "one"); + assert_eq!(result["model"], "provider-model"); + assert_eq!(result["usage_info"]["pages_processed"], 1); + assert_eq!(result["document_annotation"]["language"], "en"); + assert_eq!(result["future"], "kept"); +} + +#[test] +fn response_codec_rejects_missing_empty_and_malformed_content() { + for value in [ + json!({"choices":[]}), + json!({"choices":[{"message":{"content":""}}]}), + json!({"choices":[{"message":{"content":"{\"pages\":[{\"markdown\":42}]}"}}]}), + json!({"choices":[{"message":{"content":{"pages":[{"markdown":42}]}}}]}), + ] { + let result = serde_json::from_value::(value) + .map_err(|_| ()) + .and_then(|response| transform_ocr_response("model", response).map_err(|_| ())); + assert!(result.is_err()); + } +} diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index d1828dfb816..cecd8869741 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -21,6 +21,7 @@ fn request_boundary_selects_mistral_and_rejects_unknown_providers() { .as_object() .unwrap() .clone(), + input_sources: Default::default(), timeout_seconds: None, }; assert!(decode_request(request).is_ok()); @@ -33,6 +34,7 @@ fn request_boundary_selects_mistral_and_rejects_unknown_providers() { custom_llm_provider: Some("unknown".into()), extra_headers: None, optional_params: serde_json::Map::new(), + input_sources: Default::default(), timeout_seconds: None, }) .is_err() diff --git a/litellm-rust/crates/core/tests/ocr/support.rs b/litellm-rust/crates/core/tests/ocr/support.rs index d9720019419..a2e67dffc7d 100644 --- a/litellm-rust/crates/core/tests/ocr/support.rs +++ b/litellm-rust/crates/core/tests/ocr/support.rs @@ -8,7 +8,11 @@ use crate::ocr::wire::{OcrWireRequest, decode_request}; use crate::ocr::{LiteLLMOcrRequest, LiteLLMOcrResponse, OcrClient}; pub(crate) fn ocr_client() -> OcrClient { - OcrClient::for_test(reqwest::Client::new()) + let document_http = reqwest::Client::builder() + .redirect(reqwest::redirect::Policy::none()) + .build() + .expect("test document client builds"); + OcrClient::for_test(reqwest::Client::new(), document_http) } pub(crate) async fn perform_ocr( @@ -26,6 +30,7 @@ pub(crate) fn wire_request(model: &str, base: &str, options: Value) -> LiteLLMOc custom_llm_provider: None, extra_headers: None, optional_params: options.as_object().unwrap().clone(), + input_sources: Default::default(), timeout_seconds: Some(2.0), }) .unwrap() diff --git a/litellm-rust/crates/core/tests/reducto_ocr.rs b/litellm-rust/crates/core/tests/reducto_ocr.rs new file mode 100644 index 00000000000..8e86e4713ef --- /dev/null +++ b/litellm-rust/crates/core/tests/reducto_ocr.rs @@ -0,0 +1,232 @@ +use std::sync::Arc; + +use rstest::rstest; +use serde_json::{Value, json}; + +use super::hooks::{OcrDuringCallRequest, OcrHookFuture, OcrHooks}; +use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; + +fn request_body(request: &str) -> Value { + serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap() +} + +#[rstest] +#[case( + "reducto/parse-v3", + json!({ + "formatting":{"table_output_format":"html"}, + "retrieval":{"chunk_mode":"section"}, + "settings":{"ocr_system":"standard"}, + "future_ocr_option":true, + "extra_body":{"provider_option":"value"} + }), + "reducto://already.pdf", + json!({ + "input":"reducto://already.pdf", + "formatting":{"table_output_format":"html"}, + "retrieval":{"chunk_mode":"section"}, + "settings":{"ocr_system":"standard"}, + "future_ocr_option":true, + "provider_option":"value" + }) +)] +#[case( + "reducto/parse-legacy", + json!({ + "enhance":{"agentic":[{"type":"table"}]}, + "future_ocr_option":true, + "extra_body":{"provider_option":"value"} + }), + "reducto://legacy.pdf", + json!({ + "document_url":"reducto://legacy.pdf", + "options":{"enhance":{"agentic":[{"type":"table"}]}}, + "future_ocr_option":true, + "provider_option":"value" + }) +)] +#[tokio::test] +async fn request_mapping_matches_python( + #[case] model: &str, + #[case] options: Value, + #[case] source: &str, + #[case] expected: Value, +) { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ + "result":{"chunks":[]} + }))]) + .await; + let mut request = wire_request(model, &base, options); + request.document = request.document.with_source(source.into()); + + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert!(requests[0].starts_with("POST /parse ")); + assert_eq!(request_body(&requests[0]), expected); +} + +#[rstest] +#[case("parse-v3")] +#[case("parse-legacy")] +#[tokio::test] +async fn data_uri_upload_preserves_multipart_headers(#[case] model: &str) { + let (base, seen, server) = mock_server(vec![ + MockResponse::json(json!({"file_id":"reducto://uploaded.pdf"})), + MockResponse::json(json!({"result":{"chunks":[{"content":"hello"}]}})), + ]) + .await; + let mut request = wire_request(&format!("reducto/{model}"), &base, json!({})); + request.connection.extra_headers = vec![ + ("Content-Type".into(), "application/json".into()), + ("X-Trace".into(), "upload-test".into()), + ]; + + let response = perform_ocr(request).await.unwrap(); + server.await.unwrap(); + assert_eq!(response.pages[0]["markdown"], "hello"); + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 2); + assert!(requests[0].starts_with("POST /upload ")); + assert!( + requests[0] + .to_ascii_lowercase() + .contains("content-type: multipart/form-data; boundary=") + ); + assert!(requests[0].contains("x-trace: upload-test")); + assert!(requests[0].contains("application/pdf")); + assert!(requests[0].contains("abc")); + assert!(requests[1].starts_with("POST /parse ")); +} + +#[rstest] +#[case(json!({"file_id":""}))] +#[case(json!({}))] +#[case(json!({"file_id":null}))] +#[tokio::test] +async fn invalid_upload_ids_stop_before_parse(#[case] response: Value) { + let (base, seen, server) = mock_server(vec![MockResponse::json(response)]).await; + let error = perform_ocr(wire_request("reducto/parse-v3", &base, json!({}))) + .await + .unwrap_err(); + server.await.unwrap(); + assert!(error.to_string().contains("file_id")); + assert_eq!(seen.lock().unwrap().len(), 1); +} + +#[tokio::test] +async fn upload_failure_stops_before_parse() { + let (base, seen, server) = mock_server(vec![MockResponse { + status: 503, + headers: vec![], + body: json!({"error":"unavailable"}), + }]) + .await; + assert!( + perform_ocr(wire_request("reducto/parse-v3", &base, json!({}))) + .await + .is_err() + ); + server.await.unwrap(); + assert_eq!(seen.lock().unwrap().len(), 1); +} + +#[rstest] +#[case("https://example.com/a.pdf")] +#[case("reducto://")] +#[case("data:application/pdf;base64")] +#[case("data:application/pdf;base64,INVALID!")] +#[tokio::test] +async fn rejects_invalid_document_sources_before_network(#[case] source: &str) { + let mut request = wire_request("reducto/parse-v3", "http://127.0.0.1:1", json!({})); + request.document = request.document.with_source(source.into()); + assert!(perform_ocr(request).await.is_err()); +} + +#[test] +fn response_normalization_groups_blocks_and_distinguishes_null_result() { + use crate::ocr::codecs::reducto::{ReductoResponse, transform_ocr_response}; + + let raw = json!({"usage":{"num_pages":"2","credits":"3"},"result":{"chunks":[ + {"blocks":[{"content":"B","bbox":{"page":2},"kind":"table"}]}, + {"blocks":[{"content":"A","bbox":{"page":1},"kind":"text"},{"content":"C","bbox":{"page":1}}]} + ]}}); + let response: ReductoResponse = serde_json::from_value(raw).unwrap(); + let normalized = transform_ocr_response("parse-v3", response) + .unwrap() + .into_json(); + assert_eq!(normalized["pages"][0]["markdown"], "A\n\nC"); + assert_eq!(normalized["pages"][1]["markdown"], "B"); + assert_eq!(normalized["pages"][1]["blocks"][0]["kind"], "table"); + assert_eq!(normalized["usage_info"]["pages_processed"], 2); + assert_eq!(normalized["usage_info"]["credits"], 3.0); + + let missing: ReductoResponse = + serde_json::from_value(json!({"chunks":[{"content":"text"}]})).unwrap(); + let missing = transform_ocr_response("parse-v3", missing).unwrap(); + assert_eq!(missing.pages[0]["markdown"], "text"); + let null: ReductoResponse = serde_json::from_value( + json!({"result":null,"chunks":[{"content":"ignored"}],"usage":null}), + ) + .unwrap(); + let null = transform_ocr_response("parse-v3", null).unwrap(); + assert!(null.pages.is_empty()); +} + +#[tokio::test] +async fn facade_omits_native_response_by_default_and_preserves_auth_priority() { + let raw = json!({"job_id":"job-1","result":{"chunks":[]}}); + let (base, seen, server) = mock_server(vec![MockResponse::json(raw)]).await; + let mut request = wire_request("reducto/parse-v3", &base, json!({})); + request.document = request.document.with_source("reducto://ready.pdf".into()); + request.connection.extra_headers = vec![("authorization".into(), "Bearer existing".into())]; + + let response = perform_ocr(request).await.unwrap(); + server.await.unwrap(); + assert_eq!(response.provider_native_response, None); + assert!( + seen.lock().unwrap()[0] + .to_ascii_lowercase() + .contains("authorization: bearer existing") + ); +} + +struct RewriteDocument; + +impl OcrHooks for RewriteDocument { + fn has_guardrails(&self) -> bool { + true + } + + fn during_call( + &self, + request: OcrDuringCallRequest, + ) -> OcrHookFuture<'_, OcrDuringCallRequest> { + Box::pin(async move { + assert_eq!( + request.body["document_url"], + "data:application/pdf;base64,YWJj" + ); + Ok(OcrDuringCallRequest { + body: json!({"type":"document_url","document_url":"reducto://guarded.pdf"}), + ..request + }) + }) + } +} + +#[tokio::test] +async fn guardrail_rewrites_document_before_upload() { + let (base, seen, server) = + mock_server(vec![MockResponse::json(json!({"result":{"chunks":[]}}))]).await; + let mut request = wire_request("reducto/parse-v3", &base, json!({})); + request.hooks = Arc::new(RewriteDocument); + + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert!(requests[0].starts_with("POST /parse ")); + assert!(requests[0].contains("reducto://guarded.pdf")); +} diff --git a/litellm-rust/crates/core/tests/vertex_ai_deepseek_ocr.rs b/litellm-rust/crates/core/tests/vertex_ai_deepseek_ocr.rs new file mode 100644 index 00000000000..6d3061d8f5d --- /dev/null +++ b/litellm-rust/crates/core/tests/vertex_ai_deepseek_ocr.rs @@ -0,0 +1,83 @@ +use serde_json::{Value, json}; + +use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; +use crate::auth::InputSource; + +fn request_body(request: &str) -> Value { + serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap() +} + +#[tokio::test] +async fn facade_executes_vertex_deepseek_at_the_openai_endpoint() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ + "choices":[{"message":{"content":"recognized"}}], + "usage":{"prompt_tokens":1} + }))]) + .await; + let mut request = wire_request( + "vertex_ai/deepseek-ocr-maas", + &base, + json!({ + "vertex_project":"project-1", + "vertex_location":"europe-west4", + "temperature":0.1, + "future_ocr_option":true, + "extra_body":{"provider_option":"value"} + }), + ); + request.document = request + .document + .with_source("gs://bucket/document.pdf".into()); + + let response = perform_ocr(request).await.unwrap(); + server.await.unwrap(); + assert_eq!(response.pages[0]["markdown"], "recognized"); + assert_eq!(response.usage_info.unwrap()["prompt_tokens"], 1); + let requests = seen.lock().unwrap(); + assert!(requests[0].starts_with( + "POST /v1/projects/project-1/locations/europe-west4/endpoints/openapi/chat/completions " + )); + assert!( + requests[0] + .to_ascii_lowercase() + .contains("authorization: bearer test-key") + ); + let body = request_body(&requests[0]); + assert_eq!(body["model"], "deepseek-ai/deepseek-ocr-maas"); + assert_eq!(body["temperature"], 0.1); + assert!(body.get("future_ocr_option").is_none()); + assert!(body.get("extra_body").is_none()); + assert_eq!( + body["messages"][0]["content"][0], + json!({"type":"document_url","document_url":"gs://bucket/document.pdf"}) + ); +} + +#[test] +fn host_registration_selects_deepseek_without_affecting_mistral() { + assert!(crate::ocr::wire::is_supported_request( + "deepseek-ocr-maas", + Some("vertex_ai") + )); + assert!(crate::ocr::wire::is_supported_request( + "mistral-ocr-maas", + Some("vertex_ai") + )); +} + +#[tokio::test] +async fn request_controlled_api_base_is_rejected_before_vertex_auth() { + let mut request = wire_request( + "vertex_ai/deepseek-ocr-maas", + "https://caller.example", + json!({"vertex_project":"project-1"}), + ); + request.connection.api_base_source = InputSource::Request; + + let error = perform_ocr(request).await.unwrap_err(); + assert!( + error + .to_string() + .contains("request-controlled Vertex AI endpoint") + ); +} diff --git a/litellm-rust/crates/core/tests/vertex_ai_ocr.rs b/litellm-rust/crates/core/tests/vertex_ai_ocr.rs new file mode 100644 index 00000000000..96a19dd62b4 --- /dev/null +++ b/litellm-rust/crates/core/tests/vertex_ai_ocr.rs @@ -0,0 +1,161 @@ +use serde_json::{Value, json}; + +use super::test_support::{MockResponse, mock_server, perform_ocr, wire_request}; +use crate::auth::InputSource; + +fn request_body(request: &str) -> Value { + serde_json::from_str(request.split_once("\r\n\r\n").unwrap().1).unwrap() +} + +#[tokio::test] +async fn facade_executes_vertex_mistral_with_resolved_project_and_location() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({ + "pages":[{"index":0,"markdown":"hello"}], + "usage_info":{"pages_processed":1} + }))]) + .await; + let request = wire_request( + "vertex_ai/mistral-ocr-maas", + &base, + json!({ + "vertex_project":"project-1", + "vertex_location":"europe-west4", + "extract_footer":true + }), + ); + + let response = perform_ocr(request).await.unwrap(); + server.await.unwrap(); + assert_eq!(response.pages[0]["markdown"], "hello"); + let requests = seen.lock().unwrap(); + assert_eq!(requests.len(), 1); + assert!(requests[0].starts_with( + "POST /v1/projects/project-1/locations/europe-west4/publishers/mistralai/models/mistral-ocr-maas:rawPredict " + )); + assert!( + requests[0] + .to_ascii_lowercase() + .contains("authorization: bearer test-key") + ); + assert_eq!( + request_body(&requests[0]), + json!({ + "model":"mistral-ocr-maas", + "document":{"type":"document_url","document_url":"data:application/pdf;base64,YWJj"}, + "extract_footer":true + }) + ); +} + +#[tokio::test] +async fn supplied_authorization_is_forwarded_without_a_static_token() { + let (base, seen, server) = mock_server(vec![MockResponse::json(json!({"pages":[]}))]).await; + let mut request = wire_request( + "vertex_ai/model", + &base, + json!({"vertex_project":"project-1"}), + ); + request.connection.api_key = None; + request.connection.extra_headers = vec![("authorization".into(), "Bearer supplied".into())]; + + perform_ocr(request).await.unwrap(); + server.await.unwrap(); + assert!( + seen.lock().unwrap()[0] + .to_ascii_lowercase() + .contains("authorization: bearer supplied") + ); +} + +#[tokio::test] +async fn invalid_credentials_fail_before_provider_http() { + let request = wire_request( + "vertex_ai/model", + "http://127.0.0.1:1", + json!({"vertex_credentials": true}), + ); + let error = perform_ocr(request).await.unwrap_err(); + assert!(error.to_string().contains("vertex_credentials")); +} + +#[tokio::test] +async fn request_controlled_api_base_is_rejected_before_vertex_auth() { + let mut request = wire_request( + "vertex_ai/mistral-ocr-maas", + "https://caller.example", + json!({"vertex_project":"project-1"}), + ); + request.connection.api_base_source = InputSource::Request; + + let error = perform_ocr(request).await.unwrap_err(); + assert!( + error + .to_string() + .contains("request-controlled Vertex AI endpoint") + ); +} + +#[tokio::test] +async fn adapters_build_complete_requests_and_share_mistral_normalization() { + use std::time::Duration; + + use crate::ocr::adapters::{MistralAdapter, OcrAdapter, VertexMistralAdapter}; + use crate::ocr::test_support::ocr_client; + + let client = ocr_client(); + let options = json!({ + "pages": [0, 2], + "include_image_base64": true, + "vertex_project": "project-1", + "vertex_location": "us-central1", + "unknown": "ignored" + }); + let direct = wire_request( + "mistral/mistral-ocr-maas", + "https://mistral.test", + options.clone(), + ); + let vertex = wire_request("vertex_ai/mistral-ocr-maas", "https://vertex.test", options); + let direct_http = MistralAdapter + .prepare_request(&direct, &client) + .await + .unwrap(); + let vertex_http = VertexMistralAdapter + .prepare_request(&vertex, &client) + .await + .unwrap(); + assert_eq!(direct_http.url().as_str(), "https://mistral.test/v1/ocr"); + assert_eq!( + vertex_http.url().as_str(), + "https://vertex.test/v1/projects/project-1/locations/us-central1/publishers/mistralai/models/mistral-ocr-maas:rawPredict" + ); + for http in [&direct_http, &vertex_http] { + assert_eq!(http.method(), reqwest::Method::POST); + assert_eq!(http.headers()["authorization"], "Bearer test-key"); + assert_eq!(http.headers()["content-type"], "application/json"); + assert_eq!(http.timeout(), Some(&Duration::from_secs(2))); + let body: Value = serde_json::from_slice(http.body().unwrap().as_bytes().unwrap()).unwrap(); + assert_eq!( + body, + json!({ + "model": "mistral-ocr-maas", + "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + "pages": [0, 2], + "include_image_base64": true + }) + ); + } + let payload = json!({"pages": [{"index": 0, "markdown": "hello"}], "extra": "preserved"}); + let direct_response = MistralAdapter + .transform_ocr_response(&direct, serde_json::from_value(payload.clone()).unwrap()) + .unwrap() + .into_json(); + let vertex_response = VertexMistralAdapter + .transform_ocr_response(&vertex, serde_json::from_value(payload).unwrap()) + .unwrap() + .into_json(); + assert_eq!(direct_response, vertex_response); + assert_eq!(direct_response["model"], "mistral-ocr-maas"); + assert_eq!(direct_response["object"], "ocr"); + assert_eq!(direct_response["extra"], "preserved"); +} diff --git a/litellm-rust/crates/python-bridge/src/errors.rs b/litellm-rust/crates/python-bridge/src/errors.rs index 407475dedce..e1f458ea0bc 100644 --- a/litellm-rust/crates/python-bridge/src/errors.rs +++ b/litellm-rust/crates/python-bridge/src/errors.rs @@ -42,6 +42,9 @@ pub(crate) fn chat_completions_error_to_pyerr(err: Error) -> PyErr { | Error::InvalidType { .. } | Error::MissingField(_) | Error::MissingApiKey { .. } + | Error::MissingAzureAiCredentials + | Error::MissingAzureDocumentIntelligenceCredentials + | Error::MissingReductoApiKey | Error::Routing(_) // Nothing reached the provider, so serving it on Python cannot double // bill and is the only way the caller gets an answer at all. diff --git a/litellm-rust/crates/python-bridge/src/routes/definition.rs b/litellm-rust/crates/python-bridge/src/routes/definition.rs index bc51647cbad..97313651011 100644 --- a/litellm-rust/crates/python-bridge/src/routes/definition.rs +++ b/litellm-rust/crates/python-bridge/src/routes/definition.rs @@ -225,7 +225,7 @@ mod tests { ( "ocr", "aocr", - "(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None)", + "(model, document, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, input_sources=None, timeout_seconds=None)", ), ( "transcription", diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr.rs b/litellm-rust/crates/python-bridge/src/routes/ocr.rs index cc2f8e43cea..c5def64c2f1 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr.rs @@ -2,6 +2,7 @@ use litellm_core::Error; use std::future::Future; use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; +use litellm_core::ocr::wire::{OcrWireRequest, decode_request, is_supported_request}; use pyo3::prelude::*; use serde_json::Value; @@ -21,6 +22,12 @@ fn prepare_ocr( timeout_seconds: inputs.timeout_seconds, })?; let optional_params = object_or_empty("optional_params", inputs.optional_params)?; + let input_sources = inputs + .input_sources + .map(serde_json::from_value) + .transpose() + .map_err(|error| pyo3::exceptions::PyValueError::new_err(error.to_string()))? + .unwrap_or_default(); Ok(async move { let RouteOptions { @@ -31,6 +38,22 @@ fn prepare_ocr( extra_headers, timeout, } = options; + if is_supported_request(&model, custom_llm_provider.as_deref()) { + let request = decode_request(OcrWireRequest { + model, + document, + api_key, + api_base, + custom_llm_provider, + extra_headers, + optional_params, + input_sources, + timeout_seconds: timeout.map(|value| value.as_secs_f64()), + })?; + return litellm_core::ocr::ocr(request) + .await + .map(|response| response.into_json()); + } run_ocr(OcrRequest { model: &model, document, @@ -66,8 +89,29 @@ bridge_route! { extra_headers: Option, #[pyo3(from_py_with = litellm_python_interop::from_py)] optional_params: Option, + #[pyo3(from_py_with = litellm_python_interop::from_py)] + input_sources: Option, timeout_seconds: Option, }, prepare = prepare_ocr, errors = ocr_error_to_pyerr, } + +#[cfg(test)] +mod tests { + use litellm_core::ocr::wire::is_supported_request; + + #[test] + fn native_activation_includes_migrated_providers() { + assert!(is_supported_request("model", Some("mistral"))); + assert!(is_supported_request("pixtral-12b", Some("azure_ai"))); + assert!(is_supported_request( + "documentintelligence/prebuilt-read", + Some("azure_ai") + )); + assert!(is_supported_request("parse-v3", Some("reducto"))); + assert!(is_supported_request("parse-legacy", Some("reducto"))); + assert!(is_supported_request("mistral-ocr", Some("vertex_ai"))); + assert!(is_supported_request("deepseek-ocr", Some("vertex_ai"))); + } +} diff --git a/litellm/__init__.py b/litellm/__init__.py index 157aed12a83..1dfd146a00e 100644 --- a/litellm/__init__.py +++ b/litellm/__init__.py @@ -1370,6 +1370,7 @@ from .exceptions import ( InvalidRequestError, BadRequestError, ImageFetchError, + VectorStoreSearchError, NotFoundError, PermissionDeniedError, RateLimitError, @@ -1471,9 +1472,11 @@ from .vector_stores.vector_store_registry import ( VectorStoreRegistry, VectorStoreIndexRegistry, ) +from .types.vector_stores import VectorStoreSearchFailureMode vector_store_registry: Optional[VectorStoreRegistry] = None vector_store_index_registry: Optional[VectorStoreIndexRegistry] = None +vector_store_search_failure_mode: VectorStoreSearchFailureMode = "annotate" ### RAG ### from . import rag diff --git a/litellm/caching/dual_cache.py b/litellm/caching/dual_cache.py index be761e1258b..81e2af45686 100644 --- a/litellm/caching/dual_cache.py +++ b/litellm/caching/dual_cache.py @@ -22,7 +22,7 @@ from litellm._logging import print_verbose, verbose_logger from litellm.constants import DEFAULT_MAX_REDIS_BATCH_CACHE_SIZE from .base_cache import BaseCache -from .in_memory_cache import InMemoryCache +from .in_memory_cache import DEFAULT_MAX_SIZE_IN_MEMORY, InMemoryCache from .redis_cache import RedisCache, RedisCircuitBreakerOpenError, log_redis_failure if TYPE_CHECKING: @@ -83,6 +83,9 @@ class DualCache(BaseCache): if default_redis_ttl is not None: self.default_redis_ttl = default_redis_ttl + def update_in_memory_max_size(self, max_size: int | None) -> None: + self.in_memory_cache.max_size_in_memory = DEFAULT_MAX_SIZE_IN_MEMORY if max_size is None else max_size + def attach_redis_cache( self, redis_cache: RedisCache | None = None, @@ -376,7 +379,9 @@ class DualCache(BaseCache): ) # async_batch_set_cache - async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs): + async def async_set_cache_pipeline( + self, cache_list: Sequence[tuple[str, object]], local_only: bool = False, **kwargs + ): """ Batch write values to the cache """ diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 38a9966f9f9..4058a3d72dd 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -24,11 +24,13 @@ from litellm.constants import MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB from .base_cache import BaseCache +DEFAULT_MAX_SIZE_IN_MEMORY: Final = 200 + class InMemoryCache(BaseCache): def __init__( self, - max_size_in_memory: int | None = 200, + max_size_in_memory: int | None = DEFAULT_MAX_SIZE_IN_MEMORY, default_ttl: int | None = 600, # default ttl is 10 minutes. At maximum litellm rate limiting logic requires objects to be in memory for 1 minute max_size_per_item: int | None = 1024, # 1MB = 1024KB @@ -37,7 +39,7 @@ class InMemoryCache(BaseCache): max_size_in_memory [int]: Maximum number of items in cache. done to prevent memory leaks. Use 200 items as a default """ self.max_size_in_memory = ( - max_size_in_memory if max_size_in_memory is not None else 200 + max_size_in_memory if max_size_in_memory is not None else DEFAULT_MAX_SIZE_IN_MEMORY ) # set an upper bound of 200 items in-memory self.default_ttl = default_ttl or 600 self.max_size_per_item = max_size_per_item or MAX_SIZE_PER_ITEM_IN_MEMORY_CACHE_IN_KB # 1MB = 1024KB diff --git a/litellm/constants.py b/litellm/constants.py index 028c08a691e..6b984c2673c 100644 --- a/litellm/constants.py +++ b/litellm/constants.py @@ -2002,6 +2002,7 @@ NON_INFERENCE_CALL_TYPES: Final[frozenset[str]] = frozenset( UNKNOWN_MODEL_SPEND_LOG_MODEL: Final[str] = "unknown-model" MAX_SPEND_LOG_MODEL_NAME_LENGTH: Final[int] = 256 +MCP_SPEND_LOG_MODEL_PREFIX: Final[str] = "MCP: " # PTU reservation rollup writes rows to LiteLLM_DailyTeamSpend with this # sentinel api_key so PTU flat cost stays distinguishable from real per-request diff --git a/litellm/exceptions.py b/litellm/exceptions.py index f9215267bf3..3f22a4b2dcd 100644 --- a/litellm/exceptions.py +++ b/litellm/exceptions.py @@ -10,12 +10,14 @@ ## LiteLLM versions of the OpenAI Exception Types import enum +from collections.abc import Sequence from typing import Any, Final import httpx import openai from litellm.types.utils import LiteLLMCommonStrings +from litellm.types.vector_stores import VectorStoreSearchFailure class RateLimitErrorCategory(str, enum.Enum): @@ -288,6 +290,29 @@ class ImageFetchError(BadRequestError): ) +VECTOR_STORE_SEARCH_FAILED_CODE: Final = "vector_store_search_failed" + + +class VectorStoreSearchError(BadRequestError): + def __init__( + self, + failures: Sequence[VectorStoreSearchFailure], + model: str | None = None, + llm_provider: str | None = None, + ) -> None: + self.failures: Final[tuple[VectorStoreSearchFailure, ...]] = tuple(failures) + detail: Final = "; ".join(f"{failure['vector_store_id']}: {failure['error']}" for failure in self.failures) + super().__init__( + message=( + "The request could not be grounded in every configured vector store. " + f"{len(self.failures)} vector store search(es) failed: {detail}" + ), + model=model, + llm_provider=llm_provider, + body={"type": "invalid_request_error", "code": VECTOR_STORE_SEARCH_FAILED_CODE}, + ) + + class UnprocessableEntityError(openai.UnprocessableEntityError): def __init__( self, diff --git a/litellm/experimental_mcp_client/client.py b/litellm/experimental_mcp_client/client.py index 421a3c84d75..6cecc1e6157 100644 --- a/litellm/experimental_mcp_client/client.py +++ b/litellm/experimental_mcp_client/client.py @@ -10,6 +10,7 @@ from contextlib import AbstractAsyncContextManager from datetime import timedelta from functools import partial from importlib import metadata +from types import MappingProxyType from typing import Any, Final, Protocol, TypeAlias, TypeVar import httpx @@ -77,6 +78,7 @@ from litellm._logging import verbose_logger from litellm.constants import MCP_CLIENT_TIMEOUT, MCP_NPM_CACHE_DIR, MCP_TOOL_LISTING_TIMEOUT from litellm.experimental_mcp_client.tools import list_tools_with_pagination from litellm.llms.custom_httpx.http_handler import get_ssl_configuration +from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response from litellm.types.llms.custom_http import VerifyTypes from litellm.types.mcp import ( MCPAuth, @@ -631,7 +633,9 @@ class MCPClient: auth=effective_auth, verify=ssl_config, follow_redirects=True, - event_hooks={"request": [guard]} if guard else {}, + event_hooks=MappingProxyType( + {"response": [capture_upstream_error_response], "request": [guard] if guard else []} + ), # mutable-ok: httpx types require lists of hooks ) return factory diff --git a/litellm/integrations/custom_guardrail.py b/litellm/integrations/custom_guardrail.py index 407445b2828..3abe3e1c65d 100644 --- a/litellm/integrations/custom_guardrail.py +++ b/litellm/integrations/custom_guardrail.py @@ -1144,7 +1144,7 @@ class CustomGuardrail(CustomLogger): def add_standard_logging_guardrail_information_to_request_data( self, - guardrail_json_response: Exception | str | dict | list[dict], + guardrail_json_response: object, request_data: dict, guardrail_status: GuardrailStatus, start_time: float | None = None, @@ -1289,17 +1289,10 @@ class CustomGuardrail(CustomLogger): This gets logged on downsteam Langfuse, DataDog, etc. """ - # Convert None to empty dict to satisfy type requirements - guardrail_response: dict[str, object] | str = {} if response is None else response - - # For apply_guardrail functions in custom_code_guardrail scenario, - # simplify the logged response to "allow", "deny", or "mask" - if original_inputs is not None and isinstance(response, dict): - # Check if inputs were modified by comparing them - if self._inputs_were_modified(original_inputs, response): - guardrail_response = "mask" - else: - guardrail_response = "allow" + guardrail_response: Final = self._summarize_guardrail_response( + response=response, + original_inputs=original_inputs, + ) verbose_logger.debug("Guardrail response: %s", response) @@ -1314,6 +1307,27 @@ class CustomGuardrail(CustomLogger): ) return response + def _summarize_guardrail_response( + self, + response: object, + original_inputs: Mapping[str, object] | None, + ) -> object: + """Reduce a hook's return value to what is safe to log as ``guardrail_response``. + + ``apply_guardrail`` returns the (possibly masked) inputs and ``async_pre_call_hook`` + returns the (possibly modified) request payload. Neither is a provider verdict, and + logging them verbatim ships the user's prompt to every logging sink (OTEL spans, + Datadog, spend logs), so both collapse to ``"allow"`` / ``"mask"`` by comparing + against ``original_inputs``, a copy taken before the hook ran. A string result is the + hook's own rejection message (the proxy turns it into a 400), not user input, so it is + logged as is. + """ + if response is None: + return {} + if original_inputs is None or not isinstance(response, Mapping): + return response + return "mask" if self._inputs_were_modified(original_inputs, response) else "allow" + @staticmethod def _is_guardrail_intervention(e: Exception) -> bool: """Retained spelling for existing callers; prefer ``is_guardrail_intervention``.""" @@ -1353,24 +1367,9 @@ class CustomGuardrail(CustomLogger): ) raise e - def _inputs_were_modified(self, original_inputs: dict, response: dict) -> bool: - """ - Compare original inputs with response to determine if content was modified. - - Returns True if the inputs were modified (mask scenario), False otherwise (allow scenario). - """ - # Get all keys from both dictionaries - all_keys: Final = set(original_inputs.keys()) | set(response.keys()) - - # Compare each key's value - for key in all_keys: - original_value = original_inputs.get(key) - response_value = response.get(key) - if original_value != response_value: - return True - - # No modifications detected - return False + def _inputs_were_modified(self, original_inputs: Mapping[str, object], response: Mapping[str, object]) -> bool: + """True when any baseline key's value differs in ``response`` (mask), False otherwise (allow).""" + return any(response.get(key) != value for key, value in original_inputs.items()) def mask_content_in_string( self, @@ -1477,6 +1476,31 @@ def _sync_guardrail_info_to_logging_obj(request_data: dict, logging_obj: object) _append_slg_to_litellm_params(mcd.get("litellm_params"), entries) +_PRE_CALL_CONTENT_KEYS: Final = frozenset( + {"messages", "input", "prompt", "system", "instructions", "tools", "functions", "function_call", "tool_choice"} +) + + +def _original_inputs_for( + func_name: str, + kwargs: Mapping[str, object], + request_data: Mapping[str, object], + event_type: GuardrailEventHooks | None, +) -> dict | None: # mutable-ok: matches _process_response(original_inputs=) signature + """Baseline the hook's return value is compared against to decide "allow" vs "mask". + + ``apply_guardrail`` masks a fresh ``inputs`` dict, so that dict is the baseline. Pre-call + hooks edit the request in place and return it, so the baseline is a deep copy of the + prompt-bearing keys taken before the hook runs. + """ + if func_name == "apply_guardrail": + inputs: Final = kwargs.get("inputs") + return inputs if isinstance(inputs, dict) else None + if event_type != GuardrailEventHooks.pre_call: + return None + return {key: copy.deepcopy(value) for key, value in request_data.items() if key in _PRE_CALL_CONTENT_KEYS} + + def log_guardrail_information(func): """ Decorator to add standard logging guardrail information to any function @@ -1535,9 +1559,7 @@ def log_guardrail_information(func): event_type: Final = _infer_event_type_from_function_name(func.__name__) # Store original inputs for comparison (for apply_guardrail functions) - original_inputs = None - if func.__name__ == "apply_guardrail" and "inputs" in kwargs: - original_inputs = kwargs.get("inputs") + original_inputs: Final = _original_inputs_for(func.__name__, kwargs, request_data, event_type) logging_obj: Final = kwargs.get("logging_obj") or request_data.get("litellm_logging_obj") self_recorded_token: Final = _guardrail_self_recorded.set(False) @@ -1577,9 +1599,7 @@ def log_guardrail_information(func): event_type: Final = _infer_event_type_from_function_name(func.__name__) # Store original inputs for comparison (for apply_guardrail functions) - original_inputs = None - if func.__name__ == "apply_guardrail" and "inputs" in kwargs: - original_inputs = kwargs.get("inputs") + original_inputs: Final = _original_inputs_for(func.__name__, kwargs, request_data, event_type) logging_obj: Final = kwargs.get("logging_obj") or request_data.get("litellm_logging_obj") self_recorded_token: Final = _guardrail_self_recorded.set(False) diff --git a/litellm/integrations/custom_logger.py b/litellm/integrations/custom_logger.py index d445a3adf14..62ca6b0254e 100644 --- a/litellm/integrations/custom_logger.py +++ b/litellm/integrations/custom_logger.py @@ -886,12 +886,16 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac return LITELLM_METADATA_FIELD return OLD_LITELLM_METADATA_FIELD + def redacts_messages_itself(self) -> bool: + return False + def redact_standard_logging_payload_from_model_call_details(self, model_call_details: dict) -> dict: """ Redacts or excludes fields from StandardLoggingPayload before callbacks receive it. This method handles two features: - 1. turn_off_message_logging: When True, redacts messages and responses + 1. turn_off_message_logging: When True, redacts messages and responses (unless the callback + redacts them itself, see `redacts_messages_itself`) 2. standard_logging_payload_excluded_fields: Removes specified fields entirely Return a modified copy of the provided logging payload. @@ -921,7 +925,7 @@ class CustomLogger: # https://docs.litellm.ai/docs/observability/custom_callbac } # Handle turn_off_message_logging - redact messages and responses (if not already excluded) - if turn_off_message_logging: + if turn_off_message_logging and not self.redacts_messages_itself(): redacted_str: Final = "redacted-by-litellm" if "messages" not in (excluded_fields or ()) and standard_logging_object_copy.get("messages") is not None: diff --git a/litellm/integrations/datadog/datadog_llm_obs.py b/litellm/integrations/datadog/datadog_llm_obs.py index 728bf41856f..c64a12c6d75 100644 --- a/litellm/integrations/datadog/datadog_llm_obs.py +++ b/litellm/integrations/datadog/datadog_llm_obs.py @@ -165,18 +165,54 @@ def _metadata_without_prompt_carriers(standard_logging_metadata: Mapping[str, An ) -def _redact_messages(messages: Sequence[Message]) -> tuple[Message, ...]: - """Each message's shape with its content replaced and tool payloads dropped; no message is invented.""" - return tuple( - { - "role": role if isinstance(role, str) and role in _SAFE_REDACTED_MESSAGE_ROLES else "", - "content": REDACTED_BY_LITELLM, - } - for message in messages - for role in (message.get("role", ""),) +def _safe_identifier(value: object) -> str: + return value if isinstance(value, str) else "" + + +def _redact_tool_call(tool_call: ToolCall) -> ToolCall: + return ToolCall( + name=_safe_identifier(tool_call.get("name")), + arguments=REDACTED_BY_LITELLM, + tool_id=_safe_identifier(tool_call.get("tool_id")), + type=_safe_identifier(tool_call.get("type")), ) +def _redact_tool_result(tool_result: ToolResult) -> ToolResult: + return ToolResult( + name=_safe_identifier(tool_result.get("name")), + result=REDACTED_BY_LITELLM, + tool_id=_safe_identifier(tool_result.get("tool_id")), + type=_safe_identifier(tool_result.get("type")), + ) + + +def _redact_message(message: Message) -> Message: + role: Final = message.get("role", "") + tool_calls: Final = message.get("tool_calls", ()) + tool_results: Final = message.get("tool_results", ()) + redacted: Final[Message] = { + "role": role if isinstance(role, str) and role in _SAFE_REDACTED_MESSAGE_ROLES else "", + "content": REDACTED_BY_LITELLM, + **({"tool_calls": tuple(_redact_tool_call(call) for call in tool_calls)} if tool_calls else {}), + **({"tool_results": tuple(_redact_tool_result(result) for result in tool_results)} if tool_results else {}), + } + return redacted + + +def _redact_messages(messages: Sequence[Message]) -> tuple[Message, ...]: + return tuple(_redact_message(message) for message in messages) + + +def _tool_output_tokens(messages: Sequence[Message], model: str) -> float | None: + results: Final = tuple( + result.get("result", "") for message in messages for result in message.get("tool_results", ()) + ) + if not results: + return None + return float(sum(litellm.token_counter(model=model, text=result) for result in results)) + + def _cost_dimension_tags( standard_logging_payload: StandardLoggingPayload, router_fields: Mapping[str, object] ) -> tuple[str, ...]: @@ -583,6 +619,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): standard_logging_payload=standard_logging_payload, call_type=standard_logging_payload.get("call_type"), ) + tool_output_tokens: Final = _tool_output_tokens(input_messages, standard_logging_payload.get("model") or "") input_meta: Final = InputMeta(messages=_redact_messages(input_messages) if redact_payload else input_messages) output_meta: Final = OutputMeta( messages=_redact_messages(output_messages) if redact_payload else output_messages @@ -618,7 +655,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): **({"tool_definitions": tool_definitions} if tool_definitions else {}), } - metrics: Final = self._assemble_metrics(standard_logging_payload) + metrics: Final = self._assemble_metrics(standard_logging_payload, tool_output_tokens) payload: Final[LLMObsPayload] = LLMObsPayload( parent_id=metadata_parent_id if metadata_parent_id else "undefined", @@ -676,6 +713,9 @@ class DataDogLLMObsLogger(CustomBatchLogger): ) return error_info + def redacts_messages_itself(self) -> bool: + return True + def _payload_logging_is_off(self, kwargs: Mapping[str, Any]) -> bool: return ( bool(self.turn_off_message_logging) @@ -683,7 +723,9 @@ class DataDogLLMObsLogger(CustomBatchLogger): or should_redact_message_logging(dict(kwargs)) ) - def _assemble_metrics(self, standard_logging_payload: StandardLoggingPayload) -> LLMMetrics: + def _assemble_metrics( + self, standard_logging_payload: StandardLoggingPayload, tool_output_tokens: float | None + ) -> LLMMetrics: """ Build the span metrics, including the prompt-cache counts LLM Obs charts cache savings from. @@ -721,6 +763,7 @@ class DataDogLLMObsLogger(CustomBatchLogger): else {} ), **({"reasoning_output_tokens": reasoning_output_tokens} if reasoning_output_tokens else {}), + **({"tool_output_tokens": tool_output_tokens} if tool_output_tokens is not None else {}), } return metrics 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 12ff38ce4ba..b2243060c6c 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 @@ -5,20 +5,29 @@ This hook is called before making an LLM request when a vector store is configur It searches the vector store for relevant context and appends it to the messages. """ -from collections.abc import Awaitable, Callable +from collections.abc import Awaitable, Callable, Mapping, Sequence from dataclasses import dataclass -from typing import TYPE_CHECKING, Any, Final, Protocol, cast +from typing import TYPE_CHECKING, Any, Final, Protocol, cast, get_args + +from pydantic import TypeAdapter, ValidationError +from typing_extensions import assert_never import litellm import litellm.vector_stores from litellm._logging import verbose_logger +from litellm.exceptions import VectorStoreSearchError from litellm.integrations.custom_logger import CustomLogger -from litellm.types.llms.openai import AllMessageValues, ChatCompletionUserMessage +from litellm.types.llms.openai import ( + AllMessageValues, + ChatCompletionUserMessage, + ResponsesAPIResponse, +) from litellm.types.prompts.init_prompts import PromptSpec from litellm.types.utils import CallTypes, StandardCallbackDynamicParams from litellm.types.vector_stores import ( LiteLLM_ManagedVectorStore, - VectorStoreResultContent, + VectorStoreSearchFailure, + VectorStoreSearchFailureMode, VectorStoreSearchResponse, VectorStoreSearchResult, ) @@ -30,6 +39,10 @@ if TYPE_CHECKING: else: LiteLLMLoggingObj = Any +SEARCH_FAILURES_FIELD: Final = "vector_store_search_failures" +_DEFAULT_FAILURE_MODE: Final[VectorStoreSearchFailureMode] = "annotate" +_FAILURE_MODE_ADAPTER: Final = TypeAdapter(VectorStoreSearchFailureMode) + class ProxyRuntime(Protocol): def llm_router(self) -> "Router | None": ... @@ -54,11 +67,31 @@ class ProxyServerRuntime: return prisma_client +@dataclass(frozen=True, slots=True) +class SearchSucceeded: + response: VectorStoreSearchResponse + + +@dataclass(frozen=True, slots=True) +class SearchFailed: + failure: VectorStoreSearchFailure + + +SearchOutcome = SearchSucceeded | SearchFailed + + +@dataclass(frozen=True, slots=True) +class VectorStoreAugmentation: + messages: tuple[AllMessageValues, ...] + search_results: tuple[VectorStoreSearchResponse, ...] + failures: tuple[VectorStoreSearchFailure, ...] + + class VectorStorePreCallHook(CustomLogger): CONTENT_PREFIX_STRING = "Context:\n\n" """ Custom logger that handles vector store searches before LLM calls. - + When a vector store is configured, this hook: 1. Extracts the query from the last user message 2. Calls litellm.vector_stores.search() to get relevant context @@ -101,100 +134,153 @@ class VectorStorePreCallHook(CustomLogger): Returns: Tuple of (model, modified_messages, non_default_params) """ + requested_vector_store_ids: Final = _requested_vector_store_ids(non_default_params) try: - # Check if vector store is configured - if litellm.vector_store_registry is None: - return model, messages, non_default_params - - prisma_client: Final = self.proxy_runtime.prisma_client() - llm_router: Final = self.proxy_runtime.llm_router() - - # Use database fallback to ensure synchronization across instances - vector_stores_to_run: list[ - LiteLLM_ManagedVectorStore - ] = await litellm.vector_store_registry.pop_vector_stores_to_run_with_db_fallback( + augmentation: VectorStoreAugmentation | None = await self._augment_messages( + messages=messages, non_default_params=non_default_params, tools=tools, - prisma_client=prisma_client, + litellm_logging_obj=litellm_logging_obj, ) - - if not vector_stores_to_run: - return model, messages, non_default_params - - # Extract the query from the last user message - query: Final = self._extract_query_from_messages(messages) - - if not query: - verbose_logger.debug("No query found in messages for vector store search") - return model, messages, non_default_params - - modified_messages: list[AllMessageValues] = messages.copy() - all_search_results: Final[list[VectorStoreSearchResponse]] = [] - - for vector_store_to_run in vector_stores_to_run: - # Get vector store id from the vector store config - vector_store_id = vector_store_to_run.get("vector_store_id", "") - custom_llm_provider = vector_store_to_run.get("custom_llm_provider") - litellm_params_for_vector_store = vector_store_to_run.get("litellm_params", {}) or {} - request_litellm_params = litellm_logging_obj.model_call_details.get("litellm_params", {}) - request_metadata = ( - request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {} - ) - if llm_router is not None: - search_function = cast( # cast-ok: normalize router search callable - Callable[..., Awaitable[VectorStoreSearchResponse]], - llm_router.avector_store_search, - ) - else: - search_function = cast( # cast-ok: normalize SDK search callable - Callable[..., Awaitable[VectorStoreSearchResponse]], - litellm.vector_stores.asearch, - ) - try: - search_response = await search_function( - **{ - "vector_store_id": vector_store_id, - "query": query, - "custom_llm_provider": custom_llm_provider, - "metadata": request_metadata, - **litellm_params_for_vector_store, - }, - ) - except Exception as search_error: - verbose_logger.warning( - "Vector store search failed for vector_store_id=%s, continuing without its context: %s", - vector_store_id, - search_error, - ) - continue - - verbose_logger.debug("search_response: %s", search_response) - - # Store search results for later use in citations - all_search_results.append(search_response) - - # Process search results and append as context - modified_messages = self._append_search_results_to_messages( - messages=modified_messages, search_response=search_response - ) - - # Get the number of results for logging - num_results = 0 - num_results = len(search_response.get("data", []) or []) - verbose_logger.debug("Vector store search completed. Added context from %s results", num_results) - - # Store search results as-is (already in OpenAI-compatible format) - if litellm_logging_obj and all_search_results: - litellm_logging_obj.model_call_details["search_results"] = all_search_results - - return model, modified_messages, non_default_params - except Exception as e: - verbose_logger.exception("Error in VectorStorePreCallHook: %s", e) - # Return original parameters on error + verbose_logger.exception( + "Error in VectorStorePreCallHook for vector_store_ids=%s: %s", + requested_vector_store_ids, + e, + ) return model, messages, non_default_params - def _extract_query_from_messages(self, messages: list[AllMessageValues]) -> str | None: + if augmentation is None: + return model, messages, non_default_params + + for detail, value in ( + ("search_results", list(augmentation.search_results)), + (SEARCH_FAILURES_FIELD, augmentation.failures), + ): + if value: + litellm_logging_obj.model_call_details[detail] = value + + if augmentation.failures: + failure_mode: Final = _configured_failure_mode() + match failure_mode: + case "error": + raise VectorStoreSearchError(failures=augmentation.failures, model=model) + case "annotate": + pass + case _: + assert_never(failure_mode) + + return model, list(augmentation.messages), non_default_params + + async def _augment_messages( + self, + messages: Sequence[AllMessageValues], + non_default_params: dict, + tools: list[dict] | None, + litellm_logging_obj: LiteLLMLoggingObj, + ) -> VectorStoreAugmentation | None: + if litellm.vector_store_registry is None: + return None + + prisma_client: Final = self.proxy_runtime.prisma_client() + llm_router: Final = self.proxy_runtime.llm_router() + + # Use database fallback to ensure synchronization across instances + vector_stores_to_run: Final[ + Sequence[LiteLLM_ManagedVectorStore] + ] = await litellm.vector_store_registry.pop_vector_stores_to_run_with_db_fallback( + non_default_params=non_default_params, + tools=tools, + prisma_client=prisma_client, + ) + + if not vector_stores_to_run: + return None + + query: Final = self._extract_query_from_messages(messages) + + if not query: + verbose_logger.debug("No query found in messages for vector store search") + return None + + request_litellm_params: Final = litellm_logging_obj.model_call_details.get("litellm_params", {}) + request_metadata: Final = ( + request_litellm_params.get("metadata", {}) if isinstance(request_litellm_params, dict) else {} + ) + search_function: Final = ( + cast( # cast-ok: normalize router search callable + Callable[..., Awaitable[VectorStoreSearchResponse]], + llm_router.avector_store_search, + ) + if llm_router is not None + else cast( # cast-ok: normalize SDK search callable + Callable[..., Awaitable[VectorStoreSearchResponse]], + litellm.vector_stores.asearch, + ) + ) + + outcomes: Final = tuple( + [ + await self._search_one( + vector_store=vector_store_to_run, + query=query, + request_metadata=request_metadata, + search_function=search_function, + ) + for vector_store_to_run in vector_stores_to_run + ] + ) + search_results: Final = tuple(outcome.response for outcome in outcomes if isinstance(outcome, SearchSucceeded)) + failures: Final = tuple(outcome.failure for outcome in outcomes if isinstance(outcome, SearchFailed)) + + return VectorStoreAugmentation( + messages=self._messages_with_context(messages=messages, search_results=search_results), + search_results=search_results, + failures=failures, + ) + + async def _search_one( + self, + vector_store: LiteLLM_ManagedVectorStore, + query: str, + request_metadata: Mapping[str, object], + search_function: Callable[..., Awaitable[VectorStoreSearchResponse]], + ) -> SearchOutcome: + vector_store_id: Final = vector_store.get("vector_store_id", "") + custom_llm_provider: Final = vector_store.get("custom_llm_provider") + litellm_params_for_vector_store: Final = vector_store.get("litellm_params", {}) or {} + try: + search_response: Final = await search_function( + **{ + "vector_store_id": vector_store_id, + "query": query, + "custom_llm_provider": custom_llm_provider, + "metadata": request_metadata, + **litellm_params_for_vector_store, + }, + ) + except Exception as search_error: + verbose_logger.warning( + "Vector store search failed for vector_store_id=%s, continuing without its context: %s", + vector_store_id, + search_error, + ) + return SearchFailed( + failure=VectorStoreSearchFailure( + vector_store_id=vector_store_id, + custom_llm_provider=custom_llm_provider, + error=str(search_error), + ) + ) + + verbose_logger.debug( + "Vector store search completed for vector_store_id=%s. Added context from %s results", + vector_store_id, + len(search_response.get("data") or ()), + ) + return SearchSucceeded(response=search_response) + + def _extract_query_from_messages(self, messages: Sequence[AllMessageValues]) -> str | None: """ Extract the query from the last user message. @@ -223,48 +309,40 @@ class VectorStorePreCallHook(CustomLogger): return None - def _append_search_results_to_messages( + def _messages_with_context( self, - messages: list[AllMessageValues], - search_response: VectorStoreSearchResponse, - ) -> list[AllMessageValues]: - """ - Append search results as context to the messages. + messages: Sequence[AllMessageValues], + search_results: Sequence[VectorStoreSearchResponse], + ) -> tuple[AllMessageValues, ...]: + context_messages: Final = tuple( + context_message + for search_response in search_results + if (context_message := self._context_message(search_response)) is not None + ) + if not context_messages: + return tuple(messages) + return (*messages[:-1], *context_messages, *messages[-1:]) - Args: - messages: Original list of messages - search_response: Response from vector store search - - Returns: - Modified list of messages with context appended - """ - search_response_data: Final[list[VectorStoreSearchResult] | None] = search_response.get("data") + def _context_message(self, search_response: VectorStoreSearchResponse) -> AllMessageValues | None: + """Build the context message for one vector store's results, or None when it returned nothing usable.""" + search_response_data: Final[Sequence[VectorStoreSearchResult] | None] = search_response.get("data") if not search_response_data: - return messages + return None - context_content = self.CONTENT_PREFIX_STRING + context_texts: Final = tuple( + content_text + for result in search_response_data + for content_item in (result.get("content") or ()) + if (content_text := content_item.get("text")) + ) + if not context_texts: + return None - for result in search_response_data: - result_content: list[VectorStoreResultContent] | None = result.get("content") - if result_content: - for content_item in result_content: - content_text: str | None = content_item.get("text") - if content_text: - context_content += content_text + "\n\n" - - # Only add context if we found any content - if context_content != "Context:\n\n": - # Create a copy of messages to avoid modifying the original - modified_messages: Final = messages.copy() - # Add context as a new message before the last user message - context_message: Final[ChatCompletionUserMessage] = { - "role": "user", - "content": context_content, - } - modified_messages.insert(-1, cast(AllMessageValues, context_message)) - return modified_messages - - return messages + context_message: Final[ChatCompletionUserMessage] = { + "role": "user", + "content": self.CONTENT_PREFIX_STRING + "".join(f"{text}\n\n" for text in context_texts), + } + return cast(AllMessageValues, context_message) async def async_post_call_success_deployment_hook( self, @@ -287,34 +365,34 @@ class VectorStorePreCallHook(CustomLogger): verbose_logger.debug("No litellm_logging_obj in request_data") return None - verbose_logger.debug("model_call_details keys: %s", list(litellm_logging_obj.model_call_details.keys())) - # Get search results from model_call_details (already in OpenAI format) - search_results: Final[list[VectorStoreSearchResponse] | None] = litellm_logging_obj.model_call_details.get( - "search_results" + search_results: Final[Sequence[VectorStoreSearchResponse] | None] = ( + litellm_logging_obj.model_call_details.get("search_results") + ) + search_failures: Final[Sequence[VectorStoreSearchFailure] | None] = ( + litellm_logging_obj.model_call_details.get(SEARCH_FAILURES_FIELD) ) - verbose_logger.debug("Search results found: %s", search_results is not None) - - if not search_results: - verbose_logger.debug("No search results found") + if not search_results and not search_failures: + verbose_logger.debug("No search results or search failures found") return None + if isinstance(response, ResponsesAPIResponse): + if search_failures: + setattr(response, SEARCH_FAILURES_FIELD, list(search_failures)) + return response + # Add search results to response object if hasattr(response, "choices") and response.choices: for choice in response.choices: if hasattr(choice, "message") and choice.message: - # Get existing provider_specific_fields or create new dict provider_fields = getattr(choice.message, "provider_specific_fields", None) or {} - - # Add search results (already in OpenAI-compatible format) - provider_fields["search_results"] = search_results - - # Set the provider_specific_fields + if search_results: + provider_fields["search_results"] = search_results + if search_failures: + provider_fields[SEARCH_FAILURES_FIELD] = search_failures setattr(choice.message, "provider_specific_fields", provider_fields) - verbose_logger.debug("Added %s search results to response", len(search_results)) - # Return modified response return response @@ -339,29 +417,24 @@ class VectorStorePreCallHook(CustomLogger): verbose_logger.debug("VectorStorePreCallHook.async_post_call_streaming_deployment_hook called") # Get search results from model_call_details (already in OpenAI format) - search_results: Final[list[VectorStoreSearchResponse] | None] = request_data.get("search_results") + search_results: Final[Sequence[VectorStoreSearchResponse] | None] = request_data.get("search_results") + search_failures: Final[Sequence[VectorStoreSearchFailure] | None] = request_data.get(SEARCH_FAILURES_FIELD) - verbose_logger.debug("Search results found for streaming chunk: %s", search_results is not None) - - if not search_results: - verbose_logger.debug("No search results found for streaming chunk") + if not search_results and not search_failures: + verbose_logger.debug("No search results or search failures found for streaming chunk") return response_chunk # Add search results to streaming chunk if hasattr(response_chunk, "choices") and response_chunk.choices: for choice in response_chunk.choices: if hasattr(choice, "delta") and choice.delta: - # Get existing provider_specific_fields or create new dict provider_fields = getattr(choice.delta, "provider_specific_fields", None) or {} - - # Add search results (already in OpenAI-compatible format) - provider_fields["search_results"] = search_results - - # Set the provider_specific_fields + if search_results: + provider_fields["search_results"] = search_results + if search_failures: + provider_fields[SEARCH_FAILURES_FIELD] = search_failures choice.delta.provider_specific_fields = provider_fields - verbose_logger.debug("Added %s search results to streaming chunk", len(search_results)) - # Return modified chunk return response_chunk @@ -369,3 +442,23 @@ class VectorStorePreCallHook(CustomLogger): verbose_logger.exception("Error adding search results to streaming chunk: %s", e) # Don't fail the request if search results fail to be added return response_chunk + + +def _requested_vector_store_ids(non_default_params: Mapping[str, object]) -> tuple[str, ...]: + requested: Final = non_default_params.get("vector_store_ids") + if not isinstance(requested, (list, tuple)): + return () + return tuple(str(vector_store_id) for vector_store_id in requested) + + +def _configured_failure_mode() -> VectorStoreSearchFailureMode: + try: + return _FAILURE_MODE_ADAPTER.validate_python(litellm.vector_store_search_failure_mode) + except ValidationError: + verbose_logger.warning( + "Unsupported vector_store_search_failure_mode=%r, falling back to %r. Supported modes: %s", + litellm.vector_store_search_failure_mode, + _DEFAULT_FAILURE_MODE, + ", ".join(get_args(VectorStoreSearchFailureMode)), + ) + return _DEFAULT_FAILURE_MODE diff --git a/litellm/litellm_core_utils/litellm_logging.py b/litellm/litellm_core_utils/litellm_logging.py index 6ee68ab21c5..498d662a906 100644 --- a/litellm/litellm_core_utils/litellm_logging.py +++ b/litellm/litellm_core_utils/litellm_logging.py @@ -14,7 +14,7 @@ from collections.abc import Awaitable, Callable, Iterator, Mapping, Sequence from datetime import datetime as dt_object from functools import lru_cache from types import MappingProxyType, TracebackType -from typing import TYPE_CHECKING, Any, Final, Literal, Optional, Union, cast +from typing import TYPE_CHECKING, Any, Final, Literal, Union, cast from httpx import Response from pydantic import BaseModel, JsonValue @@ -211,7 +211,7 @@ if TYPE_CHECKING: from litellm.integrations.otel.logger import OpenTelemetryV2 from litellm.integrations.otel.model.config import ExporterSpec, OpenTelemetryV2Config from litellm.litellm_core_utils.llm_cost_calc.utils import BilledTokenRates - from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, LoggedRelayResponse + from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector try: from litellm_enterprise.enterprise_callbacks.callback_controls import ( EnterpriseCallbackControls, @@ -2396,52 +2396,17 @@ class Logging(LiteLLMLoggingBaseClass): for scope in [key for key in spans_logged if isinstance(key, tuple) and key[-1:] == ("success",)]: del spans_logged[scope] - def _flush_passthrough_collected_chunks_helper( - self, - raw_bytes: list[bytes], - provider_config: "BasePassthroughConfig", - ) -> Optional["LoggedRelayResponse"]: - all_chunks: Final = provider_config._convert_raw_bytes_to_str_lines(raw_bytes) - complete_streaming_response: Final = provider_config.handle_logging_collected_chunks( - all_chunks=all_chunks, - litellm_logging_obj=self, - model=self.model, - custom_llm_provider=self.model_call_details.get("custom_llm_provider", ""), - endpoint=self.model_call_details.get("endpoint", ""), - ) - return complete_streaming_response - - def flush_passthrough_collected_chunks( - self, - raw_bytes: list[bytes], - provider_config: "BasePassthroughConfig", - ): + def flush_passthrough_collected_chunks(self, collector: "PassthroughStreamCollector"): """ - Flush collected chunks from the logging object - This is used to log the collected chunks once streaming is done on passthrough endpoints - - 1. Decode the raw bytes to string lines - 2. Get the complete streaming response from the provider config - 3. Log the complete streaming response (trigger success handler) - This is used for passthrough endpoints + Log the response a passthrough stream collector assembled once streaming is done (trigger success handler) """ - complete_streaming_response: Final = self._flush_passthrough_collected_chunks_helper( - raw_bytes=raw_bytes, - provider_config=provider_config, - ) + complete_streaming_response: Final = collector.build_logged_response(litellm_logging_obj=self) if complete_streaming_response is not None: self.success_handler(result=complete_streaming_response) - async def async_flush_passthrough_collected_chunks( - self, - raw_bytes: list[bytes], - provider_config: "BasePassthroughConfig", - ): - complete_streaming_response: Final = self._flush_passthrough_collected_chunks_helper( - raw_bytes=raw_bytes, - provider_config=provider_config, - ) + async def async_flush_passthrough_collected_chunks(self, collector: "PassthroughStreamCollector"): + complete_streaming_response: Final = collector.build_logged_response(litellm_logging_obj=self) if complete_streaming_response is not None: await self.async_success_handler(result=complete_streaming_response) @@ -6505,7 +6470,7 @@ def _get_traceback_str_for_error(error_str: str) -> str: from decimal import Decimal # used for unit testing -from typing import Any, Optional, Union +from typing import Any, Union def create_dummy_standard_logging_payload() -> StandardLoggingPayload: diff --git a/litellm/litellm_core_utils/private_json.py b/litellm/litellm_core_utils/private_json.py index 30f64c8fc27..4cd4a9b4f82 100644 --- a/litellm/litellm_core_utils/private_json.py +++ b/litellm/litellm_core_utils/private_json.py @@ -36,6 +36,21 @@ def stage_private_json(path: str, data: Mapping[str, object]) -> str: return tmp_path +def stage_private_bytes(path: str, data: bytes) -> str: + parent: Final = Path(path).parent + parent.mkdir(parents=True, exist_ok=True) + fd, tmp_path = tempfile.mkstemp(dir=str(parent), prefix=".tmp-") + try: + with os.fdopen(fd, "wb") as f: + f.write(data) + f.flush() + os.fsync(f.fileno()) + except BaseException: + Path(tmp_path).unlink(missing_ok=True) + raise + return tmp_path + + def commit_staged_json(staged: str, path: str) -> None: """Move a staged file into place, replacing whatever is there in one step""" try: @@ -68,3 +83,8 @@ def discard_staged_json(staged: str) -> None: def write_private_json(path: str, data: Mapping[str, object]) -> None: """Atomically write JSON to path with owner-only permissions (0600)""" commit_staged_json(stage_private_json(path, data), path) + + +def write_private_bytes(path: str, data: bytes) -> None: + """Atomically write bytes to path with owner-only permissions (0600); a reader holding the old file keeps it whole""" + commit_staged_json(stage_private_bytes(path, data), path) diff --git a/litellm/llms/anthropic/chat/handler.py b/litellm/llms/anthropic/chat/handler.py index d1fe4cadf40..82189461403 100644 --- a/litellm/llms/anthropic/chat/handler.py +++ b/litellm/llms/anthropic/chat/handler.py @@ -724,10 +724,11 @@ class ModelResponseIterator: content_block: Final = ContentBlockDelta(**chunk) thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] = [] - self.content_blocks.append(content_block) if "text" in content_block["delta"]: text = content_block["delta"]["text"] - elif "partial_json" in content_block["delta"]: + return text, tool_use, thinking_blocks, provider_specific_fields, reasoning_content + self.content_blocks.append(content_block) + if "partial_json" in content_block["delta"]: # Only emit tool calls if we're in a tool_use or server_tool_use block # web_search_tool_result blocks also have input_json_delta but should not be treated as tool calls # See: https://github.com/BerriAI/litellm/issues/17254 diff --git a/litellm/llms/base_llm/passthrough/transformation.py b/litellm/llms/base_llm/passthrough/transformation.py index 20180c5cfa2..ec938889b88 100644 --- a/litellm/llms/base_llm/passthrough/transformation.py +++ b/litellm/llms/base_llm/passthrough/transformation.py @@ -4,7 +4,7 @@ import re from abc import abstractmethod from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass -from typing import TYPE_CHECKING, Final, TypeAlias +from typing import TYPE_CHECKING, Final, Protocol, TypeAlias from pydantic import TypeAdapter, ValidationError @@ -80,6 +80,38 @@ def logged_relay_shape( return parsed +class PassthroughStreamCollector(Protocol): + """Consumes relayed stream bytes as they arrive and builds the response logged for spend tracking.""" + + def add(self, chunk: bytes) -> None: ... + + def build_logged_response(self, litellm_logging_obj: LiteLLMLoggingObj) -> LoggedRelayResponse | None: ... + + +class RawBytesStreamCollector: + def __init__( + self, provider_config: BasePassthroughConfig, model: str, custom_llm_provider: str, endpoint: str + ) -> None: + self._provider_config = provider_config + self._model = model + self._custom_llm_provider = custom_llm_provider + self._endpoint = endpoint + self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks + + def add(self, chunk: bytes) -> None: + self._raw_bytes.append(chunk) + + def build_logged_response(self, litellm_logging_obj: LiteLLMLoggingObj) -> LoggedRelayResponse | None: + all_chunks: Final = self._provider_config._convert_raw_bytes_to_str_lines(self._raw_bytes) + return self._provider_config.handle_logging_collected_chunks( + all_chunks=all_chunks, + litellm_logging_obj=litellm_logging_obj, + model=self._model, + custom_llm_provider=self._custom_llm_provider, + endpoint=self._endpoint, + ) + + class BasePassthroughConfig(BaseLLMModelInfo): @abstractmethod def is_streaming_request(self, endpoint: str, request_data: dict) -> bool: @@ -182,6 +214,13 @@ class BasePassthroughConfig(BaseLLMModelInfo): ) -> LoggedRelayResponse | None: return None + def create_stream_collector( + self, model: str, custom_llm_provider: str, endpoint: str + ) -> PassthroughStreamCollector: + return RawBytesStreamCollector( + provider_config=self, model=model, custom_llm_provider=custom_llm_provider, endpoint=endpoint + ) + def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]: """ Converts a list of raw bytes into a list of string lines, similar to aiter_lines() diff --git a/litellm/llms/base_llm/search/transformation.py b/litellm/llms/base_llm/search/transformation.py index 7668c6132d6..9eca3e69909 100644 --- a/litellm/llms/base_llm/search/transformation.py +++ b/litellm/llms/base_llm/search/transformation.py @@ -264,6 +264,13 @@ class BaseSearchConfig: """ raise NotImplementedError("transform_search_response must be implemented by provider") + def get_http_error_class(self, error: httpx.HTTPStatusError) -> Exception: + return self.get_error_class( + error_message=error.response.text, + status_code=error.response.status_code, + headers=dict(error.response.headers), # mutable-ok: provider error factories require dict headers + ) + def get_error_class( self, error_message: str, diff --git a/litellm/llms/bedrock/chat/invoke_handler.py b/litellm/llms/bedrock/chat/invoke_handler.py index 5f8a5544d65..5c489ecb360 100644 --- a/litellm/llms/bedrock/chat/invoke_handler.py +++ b/litellm/llms/bedrock/chat/invoke_handler.py @@ -490,10 +490,10 @@ class AWSEventStreamDecoder: reasoning_content: str | None = None thinking_blocks: list[ChatCompletionThinkingBlock | ChatCompletionRedactedThinkingBlock] | None = None - self.content_blocks.append(delta_obj) if "text" in delta_obj: text = delta_obj["text"] elif "toolUse" in delta_obj: + self.content_blocks.append(delta_obj) # When json_mode is True and this is the internal json_tool_call, # convert tool input to text content instead of tool call arguments if self.json_mode is True and self._current_tool_name == RESPONSE_FORMAT_TOOL_NAME: diff --git a/litellm/llms/bedrock/passthrough/transformation.py b/litellm/llms/bedrock/passthrough/transformation.py index fb8bc4f191f..6a120f41cb6 100644 --- a/litellm/llms/bedrock/passthrough/transformation.py +++ b/litellm/llms/bedrock/passthrough/transformation.py @@ -1,23 +1,129 @@ import json -from collections.abc import Mapping +from collections.abc import Callable, Mapping, Sequence from typing import TYPE_CHECKING, Final, Optional, cast import httpx from httpx import Response +from litellm._logging import verbose_logger from litellm.litellm_core_utils.litellm_logging import Logging -from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig +from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, PassthroughStreamCollector +from litellm.types.utils import ModelResponseStream from ..base_aws_llm import BaseAWSLLM from ..common_utils import BedrockError, BedrockEventStreamDecoderBase, BedrockModelInfo if TYPE_CHECKING: + from botocore.eventstream import EventStreamMessage from httpx import URL from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj + from litellm.llms.bedrock.chat.invoke_handler import AWSEventStreamDecoder from litellm.types.utils import CostResponseTypes +_TEXT_ONLY_DELTA_FIELDS: Final = frozenset({"content", "role"}) + + +def _plain_text_delta(chunk: ModelResponseStream) -> str | None: + """Return the delta text when the chunk carries nothing else that stream_chunk_builder reads.""" + if chunk.get("usage") is not None or chunk.provider_specific_fields or len(chunk.choices) != 1: + return None + choice: Final = chunk.choices[0] + if choice.finish_reason or choice.logprobs is not None: + return None + populated: Final = frozenset(key for key, value in choice.delta.model_dump().items() if value is not None) + if not populated <= _TEXT_ONLY_DELTA_FIELDS: + return None + content: Final = choice.delta.get("content") + return content if isinstance(content, str) else None + + +class _CoalescedChunks: + """Retains translated chunks with consecutive text deltas folded into one, so memory tracks the response text, + not the event count.""" + + def __init__(self) -> None: + self._chunks: list[ModelResponseStream] = [] # mutable-ok: instance accumulator for streaming chunks + self._open_text_parts: list[str] = [] # mutable-ok: text deltas pending fold into self._chunks[-1] + + def add(self, chunk: ModelResponseStream) -> None: + text: Final = _plain_text_delta(chunk) + if text is not None and self._open_text_parts: + self._open_text_parts.append(text) + return + self._seal_text_run() + self._chunks.append(chunk) + if text is not None: + self._open_text_parts.append(text) + + def _seal_text_run(self) -> None: + if len(self._open_text_parts) > 1: + self._chunks[-1].choices[0].delta.content = "".join(self._open_text_parts) + self._open_text_parts.clear() + + def chunks(self) -> Sequence[ModelResponseStream]: + self._seal_text_run() + return self._chunks + + +def _translate_message(decoder: "AWSEventStreamDecoder", message: str) -> ModelResponseStream | None: + from litellm.litellm_core_utils.streaming_handler import ( + convert_generic_chunk_to_model_response_stream, + generic_chunk_has_all_required_fields, + ) + from litellm.types.utils import GenericStreamingChunk + + translated_chunk: Final = decoder._chunk_parser(chunk_data=json.loads(message)) + if isinstance(translated_chunk, ModelResponseStream): + return translated_chunk + if generic_chunk_has_all_required_fields(cast(dict, translated_chunk)): + return convert_generic_chunk_to_model_response_stream(cast(GenericStreamingChunk, translated_chunk)) + return None + + +def _build_logged_response( + chunks: Sequence[ModelResponseStream], litellm_logging_obj: "LiteLLMLoggingObj" +) -> Optional["CostResponseTypes"]: + from litellm.main import stream_chunk_builder + + if len(chunks) == 0: + return None + return stream_chunk_builder(chunks=list(chunks), logging_obj=litellm_logging_obj) + + +class BedrockEventStreamCollector: + """Decodes and translates Bedrock event-stream frames as they are relayed instead of buffering the stream.""" + + def __init__( + self, + parse_event: Callable[["EventStreamMessage"], str | None], + decoder: Optional["AWSEventStreamDecoder"], + ) -> None: + from botocore.eventstream import EventStreamBuffer + + self._parse_event = parse_event + self._decoder = decoder + self._event_stream_buffer: Final[EventStreamBuffer] = EventStreamBuffer() + self._chunks: Final = _CoalescedChunks() + + def add(self, chunk: bytes) -> None: + if self._decoder is None: + return + self._event_stream_buffer.add_data(chunk) + for event in self._event_stream_buffer: + self._add_event(self._decoder, event) + + def _add_event(self, decoder: "AWSEventStreamDecoder", event: "EventStreamMessage") -> None: + message: Final = self._parse_event(event) + translated: Final = _translate_message(decoder, message) if message is not None else None + if translated is not None: + self._chunks.add(translated) + + def build_logged_response(self, litellm_logging_obj: "LiteLLMLoggingObj") -> Optional["CostResponseTypes"]: + return _build_logged_response(self._chunks.chunks(), litellm_logging_obj) + + class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamDecoderBase, BasePassthroughConfig): def get_error_class( self, @@ -168,87 +274,32 @@ class BedrockPassthroughConfig(BaseAWSLLM, BedrockModelInfo, BedrockEventStreamD return litellm_model_response - def _convert_raw_bytes_to_str_lines(self, raw_bytes: list[bytes]) -> list[str]: - from botocore.eventstream import EventStreamBuffer - - all_chunks: Final = [] - event_stream_buffer: Final = EventStreamBuffer() - for chunk in raw_bytes: - event_stream_buffer.add_data(chunk) - for event in event_stream_buffer: - message = self._parse_message_from_event(event) - if message is not None: - all_chunks.append(message) - - return all_chunks - - def handle_logging_collected_chunks( - self, - all_chunks: list[str], - litellm_logging_obj: "LiteLLMLoggingObj", - model: str, - custom_llm_provider: str, - endpoint: str, - ) -> Optional["CostResponseTypes"]: - """ - 1. Convert all_chunks to a ModelResponseStream - 2. combine model_response_stream to model_response - 3. Return the model_response - """ - - from litellm.litellm_core_utils.streaming_handler import ( - convert_generic_chunk_to_model_response_stream, - generic_chunk_has_all_required_fields, + def create_stream_collector( + self, model: str, custom_llm_provider: str, endpoint: str + ) -> PassthroughStreamCollector: + return BedrockEventStreamCollector( + parse_event=self._parse_message_from_event, + decoder=self._get_event_stream_decoder(model=model, endpoint=endpoint), ) + + def _get_event_stream_decoder(self, model: str, endpoint: str) -> Optional["AWSEventStreamDecoder"]: from litellm.llms.bedrock.chat import get_bedrock_event_stream_decoder from litellm.llms.bedrock.chat.invoke_transformations.base_invoke_transformation import ( AmazonInvokeConfig, ) - from litellm.main import stream_chunk_builder - from litellm.types.utils import GenericStreamingChunk, ModelResponseStream - all_translated_chunks: Final = [] if "invoke" in endpoint: invoke_provider: Final = AmazonInvokeConfig.get_bedrock_invoke_provider(model) if invoke_provider is None: - raise ValueError(f"Invalid invoke provider: {invoke_provider}, for model: {model}") - obj = get_bedrock_event_stream_decoder( - invoke_provider=invoke_provider, - model=model, - sync_stream=True, - json_mode=False, - ) - elif "converse" in endpoint: - obj = get_bedrock_event_stream_decoder( - invoke_provider=None, - model=model, - sync_stream=True, - json_mode=False, - ) - else: - return None - - for chunk in all_chunks: - message = json.loads(chunk) - translated_chunk = obj._chunk_parser(chunk_data=message) - - if isinstance(translated_chunk, dict) and generic_chunk_has_all_required_fields( - cast(dict, translated_chunk) - ): - chunk_obj = convert_generic_chunk_to_model_response_stream( - cast(GenericStreamingChunk, translated_chunk) + verbose_logger.warning( + "Bedrock passthrough spend tracking skipped: no invoke provider for model %s", model ) - elif isinstance(translated_chunk, ModelResponseStream): - chunk_obj = translated_chunk - else: - continue - - all_translated_chunks.append(chunk_obj) - - if len(all_translated_chunks) > 0: - model_response: Final = stream_chunk_builder( - chunks=all_translated_chunks, - logging_obj=litellm_logging_obj, + return None + return get_bedrock_event_stream_decoder( + invoke_provider=invoke_provider, model=model, sync_stream=True, json_mode=False + ) + if "converse" in endpoint: + return get_bedrock_event_stream_decoder( + invoke_provider=None, model=model, sync_stream=True, json_mode=False ) - return model_response return None diff --git a/litellm/llms/custom_httpx/llm_http_handler.py b/litellm/llms/custom_httpx/llm_http_handler.py index e720428847d..7109e6942d1 100644 --- a/litellm/llms/custom_httpx/llm_http_handler.py +++ b/litellm/llms/custom_httpx/llm_http_handler.py @@ -1918,6 +1918,7 @@ class BaseLLMHTTPHandler: url=complete_url, headers=signed_headers, ) + response.raise_for_status() else: # A signed body must be sent verbatim, re-serializing it would break the signature response = client.post( @@ -1927,6 +1928,8 @@ class BaseLLMHTTPHandler: json=data if signed_json_body is None else None, timeout=timeout, ) + except httpx.HTTPStatusError as e: + raise provider_config.get_http_error_class(e) except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) @@ -2019,6 +2022,7 @@ class BaseLLMHTTPHandler: url=complete_url, headers=signed_headers, ) + response.raise_for_status() else: # A signed body must be sent verbatim, re-serializing it would break the signature response = await async_httpx_client.post( @@ -2028,6 +2032,8 @@ class BaseLLMHTTPHandler: json=data if signed_json_body is None else None, timeout=timeout, ) + except httpx.HTTPStatusError as e: + raise provider_config.get_http_error_class(e) except Exception as e: raise self._handle_error(e=e, provider_config=provider_config) diff --git a/litellm/llms/oci/chat/transformation.py b/litellm/llms/oci/chat/transformation.py index 98e23a59eea..8e4e41b4ac1 100644 --- a/litellm/llms/oci/chat/transformation.py +++ b/litellm/llms/oci/chat/transformation.py @@ -745,13 +745,25 @@ class OCIStreamWrapper(CustomStreamWrapper): # single-event case (terminal chunk carries the only copy of the text). self._cohere_text_emitted = False - def chunk_creator(self, chunk: Any) -> ModelResponseStream: + def _emit_chunk(self, parsed: ModelResponseStream) -> ModelResponseStream: + for choice in parsed.choices: + if getattr(choice.delta, "tool_calls", None): + self.tool_call = True + if choice.finish_reason is not None: + self.received_finish_reason = choice.finish_reason + self.sent_last_chunk = True + return self.model_response_creator(chunk={"choices": parsed.choices}) + + def chunk_creator(self, chunk: Any) -> ModelResponseStream | None: if not isinstance(chunk, str): raise ValueError(f"Chunk is not a string: {chunk}") if not chunk.startswith("data:"): raise ValueError(f"Chunk does not start with 'data:': {chunk}") + payload: Final = chunk[5:].strip() + if payload == "[DONE]": + return None try: - dict_chunk: Final = json.loads(chunk[5:]) + dict_chunk: Final = json.loads(payload) except json.JSONDecodeError as e: raise OCIError( status_code=500, @@ -774,8 +786,8 @@ class OCIStreamWrapper(CustomStreamWrapper): if getattr(choice.delta, "content", None): self._cohere_text_emitted = True break - return result - return handle_generic_stream_chunk(dict_chunk) + return self._emit_chunk(result) + return self._emit_chunk(handle_generic_stream_chunk(dict_chunk)) __all__ = [ diff --git a/litellm/llms/tinyfish/search/transformation.py b/litellm/llms/tinyfish/search/transformation.py index b688dc2cd01..460394c6f2d 100644 --- a/litellm/llms/tinyfish/search/transformation.py +++ b/litellm/llms/tinyfish/search/transformation.py @@ -247,6 +247,13 @@ class TinyfishSearchConfig(BaseSearchConfig): hidden["additional_headers"] = process_response_headers(raw_headers) return parsed + def get_http_error_class(self, error: httpx.HTTPStatusError) -> Exception: + return self._wrap_error( + error_message=error.response.text, + status_code=error.response.status_code, + headers=dict(error.response.headers), # mutable-ok: existing error wrapper requires dict headers + ) + def _wrap_error( self, error_message: str, @@ -256,8 +263,7 @@ class TinyfishSearchConfig(BaseSearchConfig): """ Build an attributed ``BaseLLMException`` from a TinyFish error body. - Used only at the call sites we control inside - ``transform_search_response`` (non-2xx, JSONDecodeError, ValidationError). + Used for HTTP status errors and response transformation errors. Not an override of ``BaseSearchConfig.get_error_class``: that path is left to inherit from the base so it auto-picks-up any future LiteLLM improvements. Trade-off: network failures (routed through LiteLLM diff --git a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py index 69fe5678de9..d113b2b4f6b 100644 --- a/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py +++ b/litellm/llms/vertex_ai/gemini/vertex_and_google_ai_studio_gemini.py @@ -23,6 +23,7 @@ from litellm.constants import ( DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_FLASH_LITE, DEFAULT_REASONING_EFFORT_MINIMAL_THINKING_BUDGET_GEMINI_2_5_PRO, ) +from litellm.exceptions import UnsupportedParamsError from litellm.litellm_core_utils.json_fragment_accumulator import JSONFragmentAccumulator from litellm.litellm_core_utils.prompt_templates.factory import ( _encode_tool_call_id_with_signature, @@ -108,6 +109,21 @@ else: StreamingChoices = Any +SUPPORTED_REASONING_EFFORTS: Final = ("minimal", "low", "medium", "high", "none", "disable") + + +def _unsupported_reasoning_effort(reasoning_effort: str) -> UnsupportedParamsError: + return UnsupportedParamsError( + message=( + f"Invalid `reasoning_effort`: {reasoning_effort!r}. " + f"Must be one of: {', '.join(repr(effort) for effort in SUPPORTED_REASONING_EFFORTS)}. " + "To drop this param, set `litellm.drop_params = True` or pass in `(.., drop_params=True)` " + "in the request - https://docs.litellm.ai/docs/completion/drop_params" + ), + status_code=400, + ) + + class VertexAIBaseConfig: def get_mapped_special_auth_params(self) -> dict: """ @@ -842,7 +858,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): "includeThoughts": False, } else: - raise ValueError(f"Invalid reasoning effort: {reasoning_effort}") + raise _unsupported_reasoning_effort(reasoning_effort) @staticmethod def _map_reasoning_effort_to_thinking_level( @@ -890,7 +906,7 @@ class VertexGeminiConfig(VertexAIBaseConfig, BaseConfig): else: return {"thinkingLevel": "low", "includeThoughts": False} else: - raise ValueError(f"Invalid reasoning effort: {reasoning_effort}") + raise _unsupported_reasoning_effort(reasoning_effort) @staticmethod def _is_thinking_budget_zero(thinking_budget: int | None) -> bool: diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b7726290f0e..aadd9bd3028 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -13118,6 +13118,22 @@ "supports_reasoning": true, "supports_tool_choice": true }, + "cerebras/qwen-3.8-27b": { + "input_cost_per_token": 9.9e-07, + "litellm_provider": "cerebras", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1.49e-06, + "source": "https://api.cerebras.ai/public/v1/models/qwen-3.8-27b", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "chatdolphin": { "input_cost_per_token": 5e-07, "litellm_provider": "nlp_cloud", @@ -31328,8 +31344,6 @@ "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07 }, "gpt-5.5-pro": { - "cache_read_input_token_cost": 3e-06, - "cache_read_input_token_cost_above_272k_tokens": 6e-06, "input_cost_per_token": 3e-05, "input_cost_per_token_above_272k_tokens": 6e-05, "input_cost_per_token_flex": 1.5e-05, @@ -31378,8 +31392,6 @@ "supports_low_reasoning_effort": false }, "gpt-5.5-pro-2026-04-23": { - "cache_read_input_token_cost": 3e-06, - "cache_read_input_token_cost_above_272k_tokens": 6e-06, "input_cost_per_token": 3e-05, "input_cost_per_token_above_272k_tokens": 6e-05, "input_cost_per_token_flex": 1.5e-05, @@ -31532,8 +31544,6 @@ "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07 }, "gpt-5.4-pro": { - "cache_read_input_token_cost": 3e-06, - "cache_read_input_token_cost_above_272k_tokens": 6e-06, "input_cost_per_token": 3e-05, "input_cost_per_token_above_272k_tokens": 6e-05, "input_cost_per_token_flex": 1.5e-05, @@ -31583,8 +31593,6 @@ "output_cost_per_token_above_272k_tokens_flex": 0.000135 }, "gpt-5.4-pro-2026-03-05": { - "cache_read_input_token_cost": 3e-06, - "cache_read_input_token_cost_above_272k_tokens": 6e-06, "input_cost_per_token": 3e-05, "input_cost_per_token_above_272k_tokens": 6e-05, "input_cost_per_token_flex": 1.5e-05, @@ -33076,6 +33084,7 @@ "supports_tool_choice": true }, "groq/gemma-7b-it": { + "deprecation_date": "2024-12-18", "input_cost_per_token": 5e-08, "litellm_provider": "groq", "max_input_tokens": 8192, @@ -33890,6 +33899,20 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "inception/mercury-2.5": { + "input_cost_per_token": 2e-07, + "litellm_provider": "inception", + "max_input_tokens": 260000, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 7.5e-07, + "source": "https://docs.inceptionlabs.ai/get-started/models", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, "text-completion-inception/mercury-edit-2": { "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 2.5e-07, @@ -39126,19 +39149,21 @@ "supports_tool_choice": true }, "openrouter/deepseek/deepseek-chat-v3.1": { - "input_cost_per_token": 2e-07, + "input_cost_per_token": 2.5e-07, "input_cost_per_token_cache_hit": 2e-08, "litellm_provider": "openrouter", "max_input_tokens": 163840, "max_output_tokens": 163840, "max_tokens": 163840, "mode": "chat", - "output_cost_per_token": 8e-07, + "output_cost_per_token": 9.5e-07, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "cache_read_input_token_cost": 1.3e-07, + "source": "https://openrouter.ai/deepseek/deepseek-chat-v3.1" }, "openrouter/deepseek/deepseek-v3.2": { "input_cost_per_token": 2.69e-07, @@ -39201,36 +39226,56 @@ "supports_tool_choice": true }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 1.32e-06, + "input_cost_per_token": 8.59908e-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": 3.96e-06, + "output_cost_per_token": 1.719816e-06, "source": "https://openrouter.ai/deepseek/deepseek-v4-pro", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "cache_read_input_token_cost": 7.1659e-08 + }, + "openrouter/deepseek/deepseek-v4.1-flash": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "cache_read_input_token_cost": 3e-09, + "litellm_provider": "openrouter", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "source": "https://openrouter.ai/deepseek/deepseek-v4.1-flash", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": false, + "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "input_cost_per_token": 1.32e-06, + "input_cost_per_token": 5.7948e-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": 3.96e-06, + "output_cost_per_token": 1.73844e-06, "source": "https://openrouter.ai/deepseek/deepseek-v4-pro-0813", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "cache_read_input_token_cost": 1.9316e-08 }, "openrouter/google/gemini-2.0-flash-001": { "deprecation_date": "2026-06-01", @@ -40009,6 +40054,28 @@ "supports_tool_choice": true, "supports_vision": true }, + "openrouter/openai/gpt-5.6-sol-pro": { + "input_cost_per_token": 2e-06, + "output_cost_per_token": 1e-05, + "cache_read_input_token_cost": 2e-07, + "cache_creation_input_token_cost": 2.5e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-5.6-sol-pro", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/openai/gpt-oss-120b": { "input_cost_per_token": 3.7e-08, "litellm_provider": "openrouter", @@ -40137,13 +40204,13 @@ "supports_tool_choice": true }, "openrouter/qwen/qwen3-235b-a22b-2507": { - "input_cost_per_token": 8.75e-08, + "input_cost_per_token": 2.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 262144, "max_tokens": 262144, "mode": "chat", - "output_cost_per_token": 3.5e-07, + "output_cost_per_token": 8.8e-07, "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-2507", "supports_function_calling": true, "supports_tool_choice": true @@ -40176,7 +40243,7 @@ "supports_vision": true }, "openrouter/qwen/qwen3.5-35b-a3b": { - "input_cost_per_token": 2.5e-07, + "input_cost_per_token": 3.125e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, @@ -40187,7 +40254,8 @@ "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "cache_read_input_token_cost": 1.5625e-07 }, "openrouter/qwen/qwen3.5-27b": { "input_cost_per_token": 1.95e-07, @@ -40204,13 +40272,13 @@ "supports_vision": true }, "openrouter/qwen/qwen3.5-122b-a10b": { - "input_cost_per_token": 2.9e-07, + "input_cost_per_token": 2.6e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 2.4e-06, + "output_cost_per_token": 2.08e-06, "source": "https://openrouter.ai/qwen/qwen3.5-122b-a10b", "supports_function_calling": true, "supports_reasoning": true, @@ -40297,18 +40365,19 @@ "supports_web_search": true }, "openrouter/z-ai/glm-4.6": { - "input_cost_per_token": 5.5e-07, + "input_cost_per_token": 4.3e-07, "litellm_provider": "openrouter", "max_input_tokens": 202800, "max_output_tokens": 131000, "max_tokens": 131000, "mode": "chat", - "output_cost_per_token": 2.2e-06, + "output_cost_per_token": 1.75e-06, "source": "https://openrouter.ai/z-ai/glm-4.6", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "cache_read_input_token_cost": 8e-08 }, "openrouter/z-ai/glm-4.6:exacto": { "input_cost_per_token": 4.5e-07, @@ -43133,7 +43202,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 1.2e-06, + "max_input_tokens": 131072, + "source": "https://api.together.xyz/v1/models" }, "together_ai/Qwen/Qwen2.5-7B-Instruct-Turbo": { "litellm_provider": "together_ai", @@ -43141,7 +43214,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "input_cost_per_token": 3e-07, + "output_cost_per_token": 3e-07, + "max_input_tokens": 32768, + "source": "https://api.together.xyz/v1/models" }, "together_ai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput": { "deprecation_date": "2026-07-10", @@ -43353,7 +43430,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "input_cost_per_token": 2e-07, + "output_cost_per_token": 2e-07, + "max_input_tokens": 32768, + "source": "https://api.together.xyz/v1/models" }, "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": { "deprecation_date": "2026-04-02", @@ -43361,7 +43442,11 @@ "mode": "chat", "supports_function_calling": true, "supports_parallel_function_calling": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 3e-07, + "max_input_tokens": 32768, + "source": "https://api.together.xyz/v1/models" }, "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": { "deprecation_date": "2026-04-16", @@ -48092,27 +48177,29 @@ "supports_tool_choice": true }, "vertex_ai/mistral-small-2503": { - "input_cost_per_token": 1e-06, + "input_cost_per_token": 1e-07, "litellm_provider": "vertex_ai-mistral_models", "max_input_tokens": 128000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3e-06, + "output_cost_per_token": 3e-07, "supports_function_calling": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing" }, "vertex_ai/mistral-small-2503@001": { - "input_cost_per_token": 1e-06, + "input_cost_per_token": 1e-07, "litellm_provider": "vertex_ai-mistral_models", "max_input_tokens": 32000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_token": 3e-06, + "output_cost_per_token": 3e-07, "supports_function_calling": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing" }, "vertex_ai/mistral-ocr-2505": { "litellm_provider": "vertex_ai", @@ -48162,15 +48249,16 @@ "supports_reasoning": true }, "vertex_ai/openai/gpt-oss-20b-maas": { - "input_cost_per_token": 7.5e-08, + "input_cost_per_token": 7e-08, "litellm_provider": "vertex_ai-openai_models", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 3e-07, - "source": "https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/gpt-oss-120b-maas", - "supports_reasoning": true + "output_cost_per_token": 2.5e-07, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "supports_reasoning": true, + "cache_read_input_token_cost": 7e-09 }, "vertex_ai/xai/grok-4.1-fast-non-reasoning": { "cache_read_input_token_cost": 5e-08, @@ -55272,11 +55360,14 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "bedrock_mantle/openai.gpt-5.6-terra": { @@ -55311,11 +55402,14 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "bedrock_mantle/openai.gpt-5.6-cyber": { @@ -55340,10 +55434,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "bedrock_mantle/openai.gpt-daybreak-blue-5.6-sol": { @@ -55372,10 +55468,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-daybreak-blue-56-sol.html" }, @@ -55411,11 +55509,14 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "us.openai.gpt-5.6-sol": { @@ -55440,8 +55541,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "global.openai.gpt-5.6-sol": { @@ -55466,8 +55570,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "us.openai.gpt-5.6-terra": { @@ -55492,8 +55599,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "global.openai.gpt-5.6-terra": { @@ -55518,8 +55628,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "us.openai.gpt-5.6-luna": { @@ -55544,8 +55657,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "global.openai.gpt-5.6-luna": { @@ -55570,8 +55686,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "bedrock_mantle/openai.gpt-6-astra": { @@ -55601,10 +55720,14 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-astra.html" }, @@ -55630,9 +55753,13 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, "supports_tool_choice": true, "supports_prompt_caching": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-astra.html" }, @@ -55658,9 +55785,13 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, "supports_tool_choice": true, "supports_prompt_caching": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-astra.html" }, @@ -55693,11 +55824,13 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "bedrock_mantle/openai.gpt-5.4": { @@ -55729,11 +55862,13 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "bedrock_mantle/google.gemma-4-31b": { @@ -56465,17 +56600,17 @@ "supports_reasoning": true, "source": "https://serverless.tensormesh.ai/v1/models/openrouter" }, - "deepseek-v4-flash": { + "deepseek-flash": { "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 1.4e-08, - "input_cost_per_token": 4.4e-07, - "input_cost_per_token_cache_hit": 1.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, "litellm_provider": "deepseek", "max_input_tokens": 1000000, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.32e-06, + "output_cost_per_token": 1.2e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", "supported_endpoints": [ "/v1/chat/completions" @@ -56489,19 +56624,45 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": false + "supports_vision": true }, - "deepseek-v4-flash-vision-exp": { + "deepseek-v4-flash": { "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 1.4e-08, - "input_cost_per_token": 4.4e-07, - "input_cost_per_token_cache_hit": 1.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, "litellm_provider": "deepseek", "max_input_tokens": 1000000, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.32e-06, + "output_cost_per_token": 1.2e-06, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "deepseek-v4-flash-vision-exp": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 393216, + "max_tokens": 393216, + "mode": "chat", + "output_cost_per_token": 1.2e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", "supported_endpoints": [ "/v1/chat/completions" @@ -56543,17 +56704,17 @@ "supports_tool_choice": true, "supports_vision": false }, - "deepseek/deepseek-v4-flash": { + "deepseek/deepseek-flash": { "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 1.4e-08, - "input_cost_per_token": 4.4e-07, - "input_cost_per_token_cache_hit": 1.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, "litellm_provider": "deepseek", "max_input_tokens": 1000000, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.32e-06, + "output_cost_per_token": 1.2e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", "supported_endpoints": [ "/v1/chat/completions" @@ -56567,19 +56728,45 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": false + "supports_vision": true }, - "deepseek/deepseek-v4-flash-vision-exp": { + "deepseek/deepseek-v4-flash": { "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 1.4e-08, - "input_cost_per_token": 4.4e-07, - "input_cost_per_token_cache_hit": 1.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, "litellm_provider": "deepseek", "max_input_tokens": 1000000, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.32e-06, + "output_cost_per_token": 1.2e-06, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "deepseek/deepseek-v4-flash-vision-exp": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 393216, + "max_tokens": 393216, + "mode": "chat", + "output_cost_per_token": 1.2e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", "supported_endpoints": [ "/v1/chat/completions" @@ -56935,6 +57122,23 @@ ], "supports_audio_input": true }, + "gpt-live-1": { + "input_cost_per_second": 0.0008333333333333334, + "litellm_provider": "openai", + "mode": "realtime", + "source": "https://developers.openai.com/api/docs/models/gpt-live-1", + "supported_modalities": [ + "text", + "audio" + ], + "supported_output_modalities": [ + "text", + "audio" + ], + "supports_audio_input": true, + "supports_audio_output": true, + "supports_function_calling": true + }, "gpt-realtime-translate": { "input_cost_per_second": 0.0005666666666666667, "litellm_provider": "openai", @@ -60039,6 +60243,7 @@ ] }, "xai/grok-imagine-image-quality": { + "deprecation_date": "2026-11-02", "input_cost_per_image": 0.05, "litellm_provider": "xai", "mode": "image_generation", @@ -60055,6 +60260,7 @@ ] }, "xai/grok-imagine-image-quality-20260403": { + "deprecation_date": "2026-11-02", "input_cost_per_image": 0.05, "litellm_provider": "xai", "mode": "image_generation", @@ -60071,6 +60277,7 @@ ] }, "xai/grok-imagine-image-quality-latest": { + "deprecation_date": "2026-11-02", "input_cost_per_image": 0.05, "litellm_provider": "xai", "mode": "image_generation", @@ -60590,6 +60797,113 @@ "output_cost_per_token": 4.7e-07, "source": "https://docs.together.ai/docs/serverless-models" }, + "together_ai/moonshotai/Kimi-K2.6": { + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 4.5e-06, + "cache_read_input_token_cost": 2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/moonshotai/Kimi-K2.5-fp4": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 2.8e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/MiniMaxAI/MiniMax-M2.7": { + "input_cost_per_token": 3e-07, + "output_cost_per_token": 1.2e-06, + "cache_read_input_token_cost": 6e-08, + "litellm_provider": "together_ai", + "max_input_tokens": 196608, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/zai-org/GLM-5": { + "input_cost_per_token": 1e-06, + "output_cost_per_token": 3.2e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 202752, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/zai-org/GLM-5.1": { + "input_cost_per_token": 1.4e-06, + "output_cost_per_token": 4.4e-06, + "cache_read_input_token_cost": 2.6e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 202752, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-R1-0528": { + "input_cost_per_token": 3e-06, + "output_cost_per_token": 7e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 163840, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/Qwen/Qwen3-Coder-Next-FP8": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 1.2e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/Qwen/Qwen3-VL-32B-Instruct": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 1.5e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/Qwen/Qwen3-VL-8B-Instruct": { + "input_cost_per_token": 1.8e-07, + "output_cost_per_token": 6.8e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/mistralai/Ministral-3-14B-Instruct-2512": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/nvidia/NVIDIA-Nemotron-Nano-9B-v2": { + "input_cost_per_token": 6e-08, + "output_cost_per_token": 2.5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 131072, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/mistralai/Mistral-7B-Instruct-v0.3": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 32768, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/Qwen/QwQ-32B": { + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 1.2e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 131072, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, "cerebras/gemma-4-31b": { "input_cost_per_token": 9.9e-07, "litellm_provider": "cerebras", @@ -61093,10 +61407,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "input_cost_per_token": 2.64e-06, "input_cost_per_token_above_272k_tokens": 5.28e-06, @@ -61126,10 +61442,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "input_cost_per_token": 2.64e-07, "input_cost_per_token_above_272k_tokens": 5.28e-07, @@ -61158,10 +61476,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "input_cost_per_token": 3.3e-06, "cache_read_input_token_cost": 3.3e-07, @@ -61320,10 +61640,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "input_cost_per_token": 3.3e-06, "cache_read_input_token_cost": 3.3e-07, @@ -61860,6 +62182,7 @@ "mode": "responses", "supports_web_search": true, "supports_function_calling": true, + "supports_reasoning": true, "input_cost_per_token": 1.15e-08, "output_cost_per_token": 1.7e-07, "cache_read_input_token_cost": 1.15e-09, @@ -61870,6 +62193,7 @@ "mode": "responses", "supports_web_search": true, "supports_function_calling": true, + "supports_reasoning": true, "input_cost_per_token": 2.5e-07, "output_cost_per_token": 2.5e-06, "cache_read_input_token_cost": 2.5e-07, @@ -62233,6 +62557,28 @@ "cache_read_input_token_cost": 2e-08, "supports_prompt_caching": true }, + "openrouter/openai/gpt-5.6-luna-pro": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 1.2e-06, + "cache_read_input_token_cost": 2e-08, + "cache_creation_input_token_cost": 2.5e-07, + "input_cost_per_token_above_272k_tokens": 4e-07, + "output_cost_per_token_above_272k_tokens": 1.8e-06, + "cache_read_input_token_cost_above_272k_tokens": 4e-08, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-5.6-luna-pro", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/openai/gpt-5.6-terra": { "input_cost_per_token": 2e-06, "output_cost_per_token": 1.2e-05, @@ -62252,6 +62598,28 @@ "cache_read_input_token_cost": 2e-07, "supports_prompt_caching": true }, + "openrouter/openai/gpt-5.6-terra-pro": { + "input_cost_per_token": 2e-06, + "output_cost_per_token": 1.2e-05, + "cache_read_input_token_cost": 2e-07, + "cache_creation_input_token_cost": 2.5e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "output_cost_per_token_above_272k_tokens": 1.8e-05, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-5.6-terra-pro", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/openai/o3": { "input_cost_per_token": 2e-06, "output_cost_per_token": 8e-06, @@ -62485,6 +62853,28 @@ "supports_pdf_input": true, "supports_prompt_caching": true }, + "openrouter/openai/gpt-6-astra-pro": { + "input_cost_per_token": 1e-05, + "output_cost_per_token": 5e-05, + "cache_read_input_token_cost": 1e-06, + "cache_creation_input_token_cost": 1.25e-05, + "input_cost_per_token_above_272k_tokens": 2e-05, + "output_cost_per_token_above_272k_tokens": 7.5e-05, + "cache_read_input_token_cost_above_272k_tokens": 2e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-6-astra-pro", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/qwen/qwen3.8-flash": { "input_cost_per_token": 1.5e-07, "output_cost_per_token": 4.7e-07, @@ -62504,9 +62894,9 @@ "supports_prompt_caching": true }, "openrouter/z-ai/glm-5.3-flash": { - "input_cost_per_token": 7.5e-08, - "output_cost_per_token": 2.5e-07, - "cache_read_input_token_cost": 1.5e-08, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "cache_read_input_token_cost": 3e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 131072, @@ -62621,6 +63011,25 @@ "supports_vision": true, "supports_prompt_caching": true }, + "openrouter/qwen/qwen3.8-max-0902": { + "input_cost_per_token": 2e-06, + "output_cost_per_token": 6e-06, + "cache_read_input_token_cost": 2.5e-07, + "cache_creation_input_token_cost": 2.5e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "source": "https://openrouter.ai/qwen/qwen3.8-max-0902", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": false, + "supports_prompt_caching": true + }, "openrouter/deepseek/deepseek-v4-flash-0731": { "input_cost_per_token": 6.5e-08, "output_cost_per_token": 1.8e-07, @@ -62692,9 +63101,9 @@ "supports_vision": false }, "openrouter/moonshotai/kimi-k3": { - "input_cost_per_token": 3e-06, - "output_cost_per_token": 1.5e-05, - "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 2.1e-06, + "output_cost_per_token": 1.053e-05, + "cache_read_input_token_cost": 2.35e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -62791,9 +63200,9 @@ "supports_prompt_caching": true }, "openrouter/z-ai/glm-5.2": { - "input_cost_per_token": 9.66e-07, - "output_cost_per_token": 3.036e-06, - "cache_read_input_token_cost": 1.932e-07, + "input_cost_per_token": 6e-07, + "output_cost_per_token": 2e-06, + "cache_read_input_token_cost": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 131072, @@ -62824,9 +63233,9 @@ "supports_vision": false }, "openrouter/moonshotai/kimi-k2.7-code": { - "input_cost_per_token": 6.6e-07, - "output_cost_per_token": 3.4e-06, - "cache_read_input_token_cost": 1.8e-07, + "input_cost_per_token": 7.1e-07, + "output_cost_per_token": 3.5e-06, + "cache_read_input_token_cost": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, @@ -63074,10 +63483,28 @@ "supports_vision": true, "supports_pdf_input": true }, + "openrouter/openai/gpt-chat-latest": { + "input_cost_per_token": 5e-06, + "output_cost_per_token": 3e-05, + "cache_read_input_token_cost": 5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 400000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-chat-latest", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": false, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/deepseek/deepseek-v4-flash": { - "input_cost_per_token": 8.778e-08, - "output_cost_per_token": 1.7556e-07, - "cache_read_input_token_cost": 1.7556e-08, + "input_cost_per_token": 8.54e-08, + "output_cost_per_token": 1.708e-07, + "cache_read_input_token_cost": 1.708e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, @@ -63110,8 +63537,8 @@ "supports_prompt_caching": true }, "openrouter/google/gemma-4-26b-a4b-it": { - "input_cost_per_token": 7e-08, - "output_cost_per_token": 3.4e-07, + "input_cost_per_token": 4.2e-08, + "output_cost_per_token": 2.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 16384, @@ -63793,7 +64220,7 @@ "supports_vision": false }, "openrouter/qwen/qwen3-next-80b-a3b-instruct": { - "input_cost_per_token": 1e-07, + "input_cost_per_token": 9e-08, "output_cost_per_token": 1.1e-06, "cache_read_input_token_cost": 7e-08, "litellm_provider": "openrouter", @@ -63919,8 +64346,8 @@ "supports_vision": false }, "openrouter/qwen/qwen3-30b-a3b-instruct-2507": { - "input_cost_per_token": 4.815e-08, - "output_cost_per_token": 1.9305e-07, + "input_cost_per_token": 9e-08, + "output_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32000, @@ -64118,8 +64545,8 @@ "supports_vision": false }, "openrouter/qwen/qwen3-14b": { - "input_cost_per_token": 1.2e-07, - "output_cost_per_token": 2.4e-07, + "input_cost_per_token": 2.275e-07, + "output_cost_per_token": 9.1e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 16384, diff --git a/litellm/models/__init__.py b/litellm/models/__init__.py index 07d1ffa743d..50eb6f8af4f 100644 --- a/litellm/models/__init__.py +++ b/litellm/models/__init__.py @@ -3,6 +3,7 @@ Domain models for LiteLLM backend. """ from litellm.models.access_group import LiteLLM_AccessGroupTable +from litellm.models.autorouter_session import LiteLLM_AutoRouterSession from litellm.models.budget import ( LiteLLM_BudgetTable, LiteLLM_BudgetTableFull, @@ -40,6 +41,7 @@ __all__ = [ "CredentialBase", "CredentialItem", "LiteLLM_AccessGroupTable", + "LiteLLM_AutoRouterSession", "LiteLLM_BudgetTable", "LiteLLM_BudgetTableFull", "LiteLLM_Config", diff --git a/litellm/models/autorouter_session.py b/litellm/models/autorouter_session.py new file mode 100644 index 00000000000..c7126236ec3 --- /dev/null +++ b/litellm/models/autorouter_session.py @@ -0,0 +1,39 @@ +""" +Auto-router per-session rollup model. + +Canonical definition for ``litellm_autoroutersession``, the row the spend flush +maintains per (api_key, session_id, router_name). +""" + +from collections.abc import Mapping +from datetime import datetime + +from litellm.types.llms.base import LiteLLMPydanticObjectBase + + +class LiteLLM_AutoRouterSession(LiteLLMPydanticObjectBase): + api_key: str + session_id: str + router_name: str + router_type: str + first_turn_at: datetime + last_turn_at: datetime + last_model: str + turns: int + spend: float + saved_spend: float + classifier_cost: float + tier_turns: Mapping[str, int] + baseline_models: Mapping[str, int] + + @property + def baseline_model(self) -> str | None: + """The baseline most of this session's turns were priced against, or None when no turn recorded one. + + A router reconfigured mid-session leaves turns priced against two baselines; the row keeps both + counts, and the label is the one that priced the most money-carrying turns rather than whatever the + router is configured with now. + """ + if not self.baseline_models: + return None + return max(self.baseline_models, key=lambda model: (self.baseline_models[model], model)) diff --git a/litellm/ocr/main.py b/litellm/ocr/main.py index df3f9d2096b..56bfd98895d 100644 --- a/litellm/ocr/main.py +++ b/litellm/ocr/main.py @@ -10,6 +10,7 @@ import re from collections.abc import Callable, Coroutine, Mapping from dataclasses import dataclass from io import IOBase +from types import MappingProxyType from typing import Any, Final, cast import httpx @@ -19,6 +20,7 @@ from litellm._logging import verbose_logger from litellm.constants import request_timeout from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj from litellm.llms.azure_ai.ocr.common_utils import ( + is_azure_cohere_parse_model, is_azure_document_intelligence_model, ) from litellm.llms.base_llm.ocr.transformation import ( @@ -52,21 +54,32 @@ class _PreparedOCRRequest: litellm_params: dict[str, object] effective_timeout: float | httpx.Timeout litellm_logging_obj: LiteLLMLoggingObj + caller_supplied_api_key: bool = True + caller_supplied_api_base: bool = True -@dataclass -class _PreparedRustOCRCall: - api_key: str | None - api_base: str | None - headers: dict[str, object] - optional_params: dict[str, object] - - -_RUST_OCR_PROVIDERS: Final = { - "mistral", - "azure_ai", - "vertex_ai", -} +_RUST_OCR_PROVIDERS: Final = frozenset({"mistral", "azure_ai", "vertex_ai"}) +_RUST_OCR_CONFIG_FIELDS: Final = frozenset( + { + "azure_ad_token", + "tenant_id", + "client_id", + "client_secret", + "azure_scope", + "azure_authority_host", + "azure_credential", + "azure_federated_token_file", + "vertex_credentials", + "vertex_ai_credentials", + "vertex_project", + "vertex_ai_project", + "vertex_location", + "vertex_ai_location", + } +) +_RUST_OCR_SECRET_FIELDS: Final = frozenset( + {"azure_ad_token", "client_secret", "azure_federated_token_file", "vertex_credentials", "vertex_ai_credentials"} +) def _prepare_ocr_request( @@ -94,6 +107,7 @@ def _prepare_ocr_request( if doc_type not in ["document_url", "image_url"]: raise ValueError(f"Invalid document type: {doc_type}. Must be 'document_url', 'image_url', or 'file'") + caller_supplied_api_key: Final = api_key is not None caller_supplied_api_base: Final = api_base is not None ( @@ -187,182 +201,256 @@ def _prepare_ocr_request( litellm_params=dict(litellm_params), effective_timeout=effective_timeout, litellm_logging_obj=litellm_logging_obj, + caller_supplied_api_key=caller_supplied_api_key, + caller_supplied_api_base=caller_supplied_api_base, ) -def _rust_ocr_supported(prepared_request: _PreparedOCRRequest) -> bool: - if prepared_request.optional_params.get(OCR_REQUEST_FORMAT_PARAM) == "native": - return False - if not prepared_request.provider_config.supports_rust_bridge(): - return False - return prepared_request.custom_llm_provider in _RUST_OCR_PROVIDERS - - -def _rust_bridge_optional_params( - prepared_request: _PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> dict[str, object]: - optional_params: Final = dict(prepared_request.optional_params) - if prepared_request.custom_llm_provider == "vertex_ai": - vertex_project: Final = ( - prepared_request.litellm_params.get("vertex_project") - or prepared_request.litellm_params.get("vertex_ai_project") - or litellm.vertex_project - or resolve_secret("VERTEXAI_PROJECT") - ) - vertex_location: Final = ( - prepared_request.litellm_params.get("vertex_location") - or prepared_request.litellm_params.get("vertex_ai_location") - or litellm.vertex_location - or resolve_secret("VERTEXAI_LOCATION") - or resolve_secret("VERTEX_LOCATION") - ) - if vertex_project is not None: - optional_params["vertex_project"] = vertex_project - if vertex_location is not None: - optional_params["vertex_location"] = vertex_location - return optional_params - - -def _rust_bridge_api_base( - prepared_request: _PreparedOCRRequest, - resolve_secret: Callable[[str], str | None], -) -> str | None: - if prepared_request.api_base is not None: - return prepared_request.api_base - if prepared_request.custom_llm_provider == "azure_ai": - if is_azure_document_intelligence_model(prepared_request.model): - return resolve_secret("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT") - return resolve_secret("AZURE_AI_API_BASE") +def _rust_ocr_provider(request: rust_ocr_bridge.LiteLLMOcrRequest) -> str | None: + if request.custom_llm_provider is not None: + return request.custom_llm_provider + prefix: Final = request.model.partition("/")[0] + if prefix in _RUST_OCR_PROVIDERS: + return prefix + if request.model.startswith("mistral-ocr"): + return "mistral" return None -def _prepare_rust_ocr_call( - prepared_request: _PreparedOCRRequest, - resolve_api_key: Callable[[str], str | None], -) -> _PreparedRustOCRCall: - provider_config: Final = prepared_request.provider_config - api_key_env_var: Final = provider_config.get_api_key_env_var() - resolved_api_key: Final = prepared_request.api_key or ( - resolve_api_key(api_key_env_var) if api_key_env_var is not None else None +def _rust_ocr_supported(request: rust_ocr_bridge.LiteLLMOcrRequest) -> bool: + provider: Final = _rust_ocr_provider(request) + if provider not in _RUST_OCR_PROVIDERS or request.kwargs.get(OCR_REQUEST_FORMAT_PARAM) == "native": + return False + if provider == "azure_ai": + return ( + not is_azure_cohere_parse_model(request.model) + and not callable(request.kwargs.get("azure_ad_token_provider")) + and request.kwargs.get("azure_username") is None + and request.kwargs.get("azure_password") is None + ) + return True + + +def _rust_bridge_optional_params( + request: rust_ocr_bridge.LiteLLMOcrRequest, + resolve_secret: Callable[[str], str | None], +) -> Mapping[str, object]: + optional_params: Final = MappingProxyType( + { + name: value + for name, value in request.kwargs.items() + if (name not in GenericLiteLLMParams.model_fields or name in _RUST_OCR_CONFIG_FIELDS) + and name not in {"litellm_logging_obj", "aocr", "litellm_call_id", "proxy_server_request"} + } ) - resolved_headers: Final = provider_config.validate_environment( - headers=prepared_request.extra_headers or {}, - model=prepared_request.model, - api_key=resolved_api_key, - api_base=prepared_request.api_base, - litellm_params=prepared_request.litellm_params, + provider: Final = _rust_ocr_provider(request) + if provider == "azure_ai" and litellm.enable_azure_ad_token_refresh is True: + return MappingProxyType({**optional_params, "enable_azure_ad_token_refresh": True}) + if provider != "vertex_ai": + return optional_params + project: Final = ( + request.kwargs.get("vertex_project") + or request.kwargs.get("vertex_ai_project") + or litellm.vertex_project + or resolve_secret("VERTEXAI_PROJECT") ) - resolved_complete_url: Final = provider_config.get_complete_url( - api_base=prepared_request.api_base, - model=prepared_request.model, - optional_params=prepared_request.optional_params, - litellm_params=prepared_request.litellm_params, + location: Final = ( + request.kwargs.get("vertex_location") + or request.kwargs.get("vertex_ai_location") + or litellm.vertex_location + or resolve_secret("VERTEXAI_LOCATION") + or resolve_secret("VERTEX_LOCATION") ) - rust_api_base: Final = _rust_bridge_api_base(prepared_request, resolve_api_key) - rust_optional_params: Final = _rust_bridge_optional_params(prepared_request, resolve_api_key) - prepared_request.litellm_logging_obj.pre_call( + credentials: Final = ( + request.kwargs.get("vertex_credentials") + or request.kwargs.get("vertex_ai_credentials") + or resolve_secret("VERTEXAI_CREDENTIALS") + ) + vertex_params: Final = MappingProxyType( + { + name: value + for name, value in ( + ("vertex_project", project), + ("vertex_location", location), + ("vertex_credentials", credentials), + ) + if value is not None + } + ) + return MappingProxyType({**optional_params, **vertex_params}) + + +def _rust_bridge_input_sources( + request: rust_ocr_bridge.LiteLLMOcrRequest, + optional_params: Mapping[str, object], +) -> Mapping[str, str]: + proxy_request: Final = request.kwargs.get("proxy_server_request") + if not isinstance(proxy_request, Mapping): + return MappingProxyType({}) + proxy_request_mapping: Final = cast( # cast-ok: runtime Mapping check loses generic key and value types + Mapping[object, object], proxy_request + ) + body_value: Final = proxy_request_mapping.get("body") + if not isinstance(body_value, Mapping): + return MappingProxyType({}) + body: Final = cast( # cast-ok: runtime Mapping check loses generic key and value types + Mapping[object, object], body_value + ) + credential_fields_value: Final = proxy_request_mapping.get("credential_fields", ()) + credential_fields: Final = ( + frozenset(name for name in credential_fields_value if isinstance(name, str)) + if isinstance(credential_fields_value, (list, tuple, set, frozenset)) + else frozenset() + ) + names: Final = frozenset(optional_params) | frozenset({"api_key", "api_base", "extra_headers"}) + request_sources: Final = MappingProxyType( + {name: "request" for name in names if name in body or name in credential_fields} + ) + if litellm.enable_azure_ad_token_refresh is True and "enable_azure_ad_token_refresh" in optional_params: + return MappingProxyType({**request_sources, "enable_azure_ad_token_refresh": "deployment"}) + return request_sources + + +def _marshal_rust_ocr_request( + request: rust_ocr_bridge.LiteLLMOcrRequest, + resolve_secret: Callable[[str], str | None], +) -> rust_ocr_bridge.LiteLLMOcrRequest: + if not isinstance(request.document, dict): + raise TypeError(f"document must be a dict with 'type' and URL/file field, got {type(request.document)}") + document: Final = ( + convert_file_document_to_url_document(request.document) + if request.document.get("type") == "file" + else request.document + ) + provider: Final = _rust_ocr_provider(request) + api_key: Final = request.api_key or resolve_secret("MISTRAL_API_KEY") if provider == "mistral" else request.api_key + optional_params: Final = _rust_bridge_optional_params(request, resolve_secret) + input_sources: Final = _rust_bridge_input_sources(request, optional_params) + logged_optional_params: Final = MappingProxyType( + {name: "****" if name in _RUST_OCR_SECRET_FIELDS else value for name, value in optional_params.items()} + ) + logged_kwargs: Final = MappingProxyType( + { + name: "****" if name in _RUST_OCR_SECRET_FIELDS else value + for name, value in request.kwargs.items() + if name != "proxy_server_request" + } + ) + logging_obj: Final = cast( # cast-ok: bridge kwargs carry the prepared logging object + LiteLLMLoggingObj, request.kwargs["litellm_logging_obj"] + ) + logging_obj.update_from_kwargs( + kwargs=dict(logged_kwargs), # mutable-ok: logging API requires an owned dict + model=request.model, + optional_params=dict(logged_optional_params), # mutable-ok: logging API requires an owned dict + litellm_params={ + "litellm_call_id": request.kwargs.get("litellm_call_id"), + "api_base": request.api_base, + }, # mutable-ok: legacy logging requires a concrete params dict + custom_llm_provider=provider, + ) + logging_obj.pre_call( input="OCR document processing", - api_key=resolved_api_key, - additional_args={ + api_key=api_key, + additional_args={ # mutable-ok: pre_call mutates the additional_args dict "complete_input_dict": { - "model": prepared_request.model, - "document": prepared_request.document, - **rust_optional_params, - }, - "api_base": resolved_complete_url, - "headers": resolved_headers, + "model": request.model, + "document": document, + **logged_optional_params, + }, # mutable-ok: callbacks consume a JSON-serializable request dict + "api_base": request.api_base or "", + "headers": request.extra_headers or {}, # mutable-ok: logging callbacks consume a concrete headers dict }, ) - return _PreparedRustOCRCall( - api_key=resolved_api_key, - api_base=rust_api_base, - headers=cast(dict[str, object], resolved_headers), - optional_params=rust_optional_params, + return rust_ocr_bridge.LiteLLMOcrRequest( + model=request.model, + document=document, + api_key=api_key, + api_base=request.api_base, + timeout=request.timeout if request.timeout is not None else request_timeout, + custom_llm_provider=request.custom_llm_provider, + extra_headers=request.extra_headers, + kwargs=optional_params, + input_sources=input_sources, ) def _map_rust_ocr_error( error: Exception, - prepared_request: _PreparedOCRRequest, + request: rust_ocr_bridge.LiteLLMOcrRequest, exception_types: tuple[type[BaseException], type[BaseException]] | None, ) -> Exception: - if exception_types is None: + if exception_types is None or not isinstance(error, exception_types[1]): return error - _, upstream_error = exception_types - if not isinstance(error, upstream_error): + provider: Final = _rust_ocr_provider(request) + if provider is None: return error - error_args: Final = cast( # cast-ok: BaseException.args is typed with Any in the standard library stubs + provider_config: Final = ProviderConfigManager.get_provider_ocr_config( + model=request.model.removeprefix(f"{provider}/"), provider=litellm.LlmProviders(provider) + ) + if provider_config is None: + return error + error_args: Final = cast( # cast-ok: Python exceptions expose positional args as a tuple tuple[object, ...], error.args ) - status_value: Final = error_args[0] if error_args else 0 - message_value: Final = error_args[1] if len(error_args) > 1 else str(error) - status: Final = status_value if isinstance(status_value, int) else 0 - message: Final = message_value if isinstance(message_value, str) else str(message_value) - error_factory: Final = cast( # cast-ok: the legacy provider interface leaves callable parameters untyped - Callable[..., Exception], prepared_request.provider_config.get_error_class + status: Final = error_args[0] if error_args and isinstance(error_args[0], int) else 500 + message: Final = str(error_args[1]) if len(error_args) > 1 else str(error) + error_factory: Final = cast( # cast-ok: provider configs expose heterogeneous exception factories + Callable[..., Exception], provider_config.get_error_class ) return error_factory( - error_message=message, - status_code=status or 500, - headers={}, # mutable-ok: provider error factories require a concrete header dict - ) + error_message=message, status_code=status or 500, headers={} + ) # mutable-ok: provider error factories require a concrete headers dict def _run_rust_ocr( - prepared_request: _PreparedOCRRequest, + request: rust_ocr_bridge.LiteLLMOcrRequest, resolve_api_key: Callable[[str], str | None], ) -> OCRResponse | None: if rust_ocr_bridge.load_rust_ocr() is None: return None - prepared: Final = _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ) + marshalled: Final = _marshal_rust_ocr_request(request, resolve_api_key) + input_sources: Final = marshalled.input_sources try: - rust_response: Final = rust_ocr_bridge.ocr( - model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout=prepared_request.effective_timeout, + response: Final = rust_ocr_bridge.ocr( + model=marshalled.model, + document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict + api_key=marshalled.api_key, + api_base=marshalled.api_base, + custom_llm_provider=marshalled.custom_llm_provider, + extra_headers=marshalled.extra_headers, + optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict + input_sources=input_sources, + timeout=marshalled.timeout, ) except Exception as error: - raise _map_rust_ocr_error(error, prepared_request, native_exception_types()) from error - if rust_response is None: - return None - return OCRResponse.model_validate(rust_response) + raise _map_rust_ocr_error(error, request, native_exception_types()) from error + return OCRResponse.model_validate(response) if response is not None else None async def _run_rust_aocr( - prepared_request: _PreparedOCRRequest, + request: rust_ocr_bridge.LiteLLMOcrRequest, resolve_api_key: Callable[[str], str | None], ) -> OCRResponse | None: if rust_ocr_bridge.load_rust_aocr() is None: return None - prepared: Final = _prepare_rust_ocr_call( - prepared_request=prepared_request, - resolve_api_key=resolve_api_key, - ) + marshalled: Final = _marshal_rust_ocr_request(request, resolve_api_key) + input_sources: Final = marshalled.input_sources try: - rust_response: Final = await rust_ocr_bridge.aocr( - model=prepared_request.model, - document=prepared_request.document, - api_key=prepared.api_key, - api_base=prepared.api_base, - custom_llm_provider=prepared_request.custom_llm_provider, - extra_headers=prepared.headers, - optional_params=prepared.optional_params, - timeout=prepared_request.effective_timeout, + response: Final = await rust_ocr_bridge.aocr( + model=marshalled.model, + document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict + api_key=marshalled.api_key, + api_base=marshalled.api_base, + custom_llm_provider=marshalled.custom_llm_provider, + extra_headers=marshalled.extra_headers, + optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict + input_sources=input_sources, + timeout=marshalled.timeout, ) except Exception as error: - raise _map_rust_ocr_error(error, prepared_request, native_exception_types()) from error - if rust_response is None: - return None - return OCRResponse.model_validate(rust_response) + raise _map_rust_ocr_error(error, request, native_exception_types()) from error + return OCRResponse.model_validate(response) if response is not None else None @client @@ -444,7 +532,29 @@ async def aocr( "extra_headers": extra_headers, "kwargs": kwargs, } + request: Final = rust_ocr_bridge.LiteLLMOcrRequest( + model=model, + document=document, + api_key=api_key, + api_base=api_base, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + kwargs=kwargs, + ) try: + if rust_enabled() and _rust_ocr_supported(request): + from litellm.secret_managers.main import get_secret_str + + rust_response: Final = await _run_rust_aocr( + request=request, + resolve_api_key=get_secret_str, + ) + if rust_response is None: + verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path") + else: + return rust_response + prepared: Final = _prepare_ocr_request( model=model, document=document, @@ -459,18 +569,6 @@ async def aocr( custom_llm_provider = prepared.custom_llm_provider completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider}) - if _rust_ocr_supported(prepared) and rust_enabled(): - from litellm.secret_managers.main import get_secret_str - - rust_response: Final = await _run_rust_aocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, - ) - if rust_response is None: - verbose_logger.debug("Async Rust OCR bridge unavailable; falling back to Python path") - else: - return rust_response - response = base_llm_http_handler.ocr( model=prepared.model, document=prepared.document, @@ -494,9 +592,11 @@ async def aocr( return response except Exception as e: + error_provider: Final = custom_llm_provider or _rust_ocr_provider(request) + error_model: Final = model.removeprefix(f"{error_provider}/") if error_provider else model raise litellm.exception_type( - model=model, - custom_llm_provider=custom_llm_provider, + model=error_model, + custom_llm_provider=error_provider, original_exception=e, completion_kwargs=completion_kwargs, extra_kwargs=kwargs, @@ -714,9 +814,31 @@ def ocr( "extra_headers": extra_headers, "kwargs": kwargs, } + request: Final = rust_ocr_bridge.LiteLLMOcrRequest( + model=model, + document=document, + api_key=api_key, + api_base=api_base, + timeout=timeout, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + kwargs=kwargs, + ) try: _is_async: Final = kwargs.pop("aocr", False) is True completion_kwargs["aocr"] = _is_async + if rust_enabled() and _rust_ocr_supported(request): + from litellm.secret_managers.main import get_secret_str + + rust_response: Final = _run_rust_ocr( + request=request, + resolve_api_key=get_secret_str, + ) + if rust_response is None: + verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path") + else: + return rust_response + prepared: Final = _prepare_ocr_request( model=model, document=document, @@ -731,18 +853,6 @@ def ocr( custom_llm_provider = prepared.custom_llm_provider completion_kwargs.update({"model": model, "custom_llm_provider": custom_llm_provider}) - if _rust_ocr_supported(prepared) and rust_enabled(): - from litellm.secret_managers.main import get_secret_str - - rust_response: Final = _run_rust_ocr( - prepared_request=prepared, - resolve_api_key=get_secret_str, - ) - if rust_response is None: - verbose_logger.debug("Rust OCR bridge unavailable; falling back to Python path") - else: - return rust_response - response: Final = base_llm_http_handler.ocr( model=prepared.model, document=prepared.document, @@ -760,9 +870,11 @@ def ocr( return response except Exception as e: + error_provider: Final = custom_llm_provider or _rust_ocr_provider(request) + error_model: Final = model.removeprefix(f"{error_provider}/") if error_provider else model raise litellm.exception_type( - model=model, - custom_llm_provider=custom_llm_provider, + model=error_model, + custom_llm_provider=error_provider, original_exception=e, completion_kwargs=completion_kwargs, extra_kwargs=kwargs, diff --git a/litellm/passthrough/main.py b/litellm/passthrough/main.py index 7076683f294..73d8bab686b 100644 --- a/litellm/passthrough/main.py +++ b/litellm/passthrough/main.py @@ -17,7 +17,7 @@ from httpx._types import CookieTypes, QueryParamTypes, RequestContent, RequestFi from litellm._logging import verbose_logger from litellm.litellm_core_utils.get_llm_provider_logic import get_llm_provider from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj -from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig +from litellm.llms.base_llm.passthrough.transformation import BasePassthroughConfig, PassthroughStreamCollector from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from litellm.passthrough.utils import CommonUtils @@ -36,6 +36,35 @@ def _as_generator(iterable: Iterator[bytes]) -> Generator[bytes, bytes, None]: yield from iterable +class _SpendCollection: + """Feeds relayed chunks to the provider's stream collector without letting spend tracking break the relay.""" + + def __init__(self, provider_config: BasePassthroughConfig, litellm_logging_obj: LiteLLMLoggingObj) -> None: + self.collector: Final[PassthroughStreamCollector] = provider_config.create_stream_collector( + model=litellm_logging_obj.model, + custom_llm_provider=litellm_logging_obj.model_call_details.get("custom_llm_provider", ""), + endpoint=litellm_logging_obj.model_call_details.get("endpoint", ""), + ) + self.chunk_count = 0 + self._failed = False + + def add(self, chunk: bytes) -> None: + self.chunk_count += 1 + if self._failed: + return + try: + self.collector.add(chunk) + except Exception as e: # noqa: BLE001 # Safe catch-all: spend tracking must never break the relayed stream + self._failed = True + verbose_logger.exception( + "Passthrough spend-tracking collector failed; spend dropped for this stream: %s", e + ) + + @property + def should_flush(self) -> bool: + return self.chunk_count > 0 and not self._failed + + class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]): def __init__( self, @@ -50,8 +79,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]): self._response: httpx.Response self._iterator: AsyncGenerator[bytes, bytes] self._litellm_logging_obj = litellm_logging_obj - self._provider_config = provider_config - self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks + self._spend = _SpendCollection(provider_config, litellm_logging_obj) self._flush_scheduled = False self._background_tasks: set[asyncio.Task] = set() # mutable-ok: instance set for background task tracking self._hidden_params: dict[str, object] = {} # mutable-ok: router attaches response headers here in place @@ -101,16 +129,13 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]): return _init().__await__() def _start_flush(self) -> None: - if self._flush_scheduled or not self._raw_bytes: + if self._flush_scheduled or not self._spend.should_flush: return self._flush_scheduled = True try: task: Final = asyncio.create_task( - self._litellm_logging_obj.async_flush_passthrough_collected_chunks( - raw_bytes=self._raw_bytes, - provider_config=self._provider_config, - ) + self._litellm_logging_obj.async_flush_passthrough_collected_chunks(collector=self._spend.collector) ) self._background_tasks.add(task) @@ -118,8 +143,8 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]): task.add_done_callback(self._background_tasks.discard) except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging verbose_logger.exception( - "Failed to schedule passthrough spend-tracking flush; %d buffered chunks dropped: %s", - len(self._raw_bytes), + "Failed to schedule passthrough spend-tracking flush; %d collected chunks dropped: %s", + self._spend.chunk_count, e, ) @@ -134,7 +159,7 @@ class AsyncPassthroughStreamingResponse(AsyncGenerator[bytes, bytes]): await self # pyright: ignore[reportGeneralTypeIssues] # structural type check misses __await__ try: chunk: Final = await anext(self._iterator) - self._raw_bytes.append(chunk) + self._spend.add(chunk) except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic self._start_flush() try: @@ -181,13 +206,12 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]): self.headers = response.headers self.status_code = response.status_code self._litellm_logging_obj = litellm_logging_obj - self._provider_config = provider_config self._iterator: Generator[bytes, bytes, None] = _as_generator(response.iter_bytes()) - self._raw_bytes: list[bytes] = [] # mutable-ok: instance buffer for streaming chunks + self._spend = _SpendCollection(provider_config, litellm_logging_obj) self._flush_scheduled = False def _start_flush(self) -> None: - if self._flush_scheduled or not self._raw_bytes: + if self._flush_scheduled or not self._spend.should_flush: return self._flush_scheduled = True @@ -195,14 +219,12 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]): try: executor.submit( - self._litellm_logging_obj.flush_passthrough_collected_chunks, - raw_bytes=self._raw_bytes, - provider_config=self._provider_config, + self._litellm_logging_obj.flush_passthrough_collected_chunks, collector=self._spend.collector ) except Exception as e: # noqa: BLE001 # Safe catch-all for verbose logging verbose_logger.exception( - "Failed to schedule passthrough spend-tracking flush; %d buffered chunks dropped: %s", - len(self._raw_bytes), + "Failed to schedule passthrough spend-tracking flush; %d collected chunks dropped: %s", + self._spend.chunk_count, e, ) @@ -212,7 +234,7 @@ class PassthroughStreamingResponse(Generator[bytes, bytes, None]): def __next__(self) -> bytes: try: chunk: Final = next(self._iterator) - self._raw_bytes.append(chunk) + self._spend.add(chunk) except Exception: # noqa: BLE001 # Safe catch-all for cleanup logic self._start_flush() try: diff --git a/litellm/proxy/_experimental/mcp_server/faults/__init__.py b/litellm/proxy/_experimental/mcp_server/faults/__init__.py index 1b9ee77d795..de7ee5c866a 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/__init__.py +++ b/litellm/proxy/_experimental/mcp_server/faults/__init__.py @@ -22,6 +22,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import ( GatewayRejected, UpstreamOAuthFault, UpstreamProtocolFault, + UpstreamRegistrationRefused, UpstreamReportedFault, ) @@ -31,6 +32,7 @@ __all__ = [ "GatewayRejected", "UpstreamOAuthFault", "UpstreamProtocolFault", + "UpstreamRegistrationRefused", "UpstreamReportedFault", "classify_upstream_dcr_rejection", "classify_upstream_token_rejection", diff --git a/litellm/proxy/_experimental/mcp_server/faults/classify.py b/litellm/proxy/_experimental/mcp_server/faults/classify.py index 2162d078c09..bb41436b495 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/classify.py +++ b/litellm/proxy/_experimental/mcp_server/faults/classify.py @@ -21,6 +21,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import ( GatewayRejected, UpstreamOAuthFault, UpstreamProtocolFault, + UpstreamRegistrationRefused, UpstreamReportedFault, ) @@ -122,11 +123,13 @@ def classify_upstream_dcr_rejection(response: httpx.Response, log_context: str) """Classify a dynamic-client-registration rejection. RFC 7591 §3.2.2 errors carry ``error`` / ``error_description`` and go through the same blame assignment as token errors (registration sends no client credentials, so credential codes stay caller-actionable); anything - without a usable ``error`` field is an upstream protocol fault.""" + without a usable ``error`` field is a registration refusal for 401/403 and a protocol fault otherwise.""" parsed: Final = _safe_json(response) fields: Final = parsed if isinstance(parsed, dict) else {} code: Final = _bounded_field(fields.get("error")) if code is None: + if response.status_code == 401 or response.status_code == 403: + return UpstreamRegistrationRefused(status_code=response.status_code) _log_out_of_contract("registration", response, log_context) return UpstreamProtocolFault(note=f"upstream registration failed with HTTP {response.status_code}") return _classify_oauth_error_code( diff --git a/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py b/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py index 64d14140a5b..3ecf5310482 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py +++ b/litellm/proxy/_experimental/mcp_server/faults/render_oauth.py @@ -11,7 +11,7 @@ from typing import Final from fastapi.responses import JSONResponse from typing_extensions import assert_never -from litellm.proxy._experimental.mcp_server.faults.types import UpstreamOAuthFault +from litellm.proxy._experimental.mcp_server.faults.types import CallerRejected, UpstreamOAuthFault from litellm.proxy._experimental.mcp_server.oauth_utils import TOKEN_NO_CACHE_HEADERS @@ -35,6 +35,24 @@ def _upstream_reported_status_and_description(code: str) -> tuple[int, str]: return 502, "the upstream authorization server reported an internal error" +def _registration_refused_description(status_code: int) -> str: + return ( + f"the upstream authorization server refused dynamic client registration (HTTP {status_code}). " + "This provider may require a pre-registered OAuth client. Configure client_id and, if required " + "by the provider, client_secret for this MCP server to skip dynamic registration" + ) + + +def _render_caller_rejected(fault: CallerRejected) -> JSONResponse: + content: Final = { + "error": fault.code, + **({"error_description": fault.description} if fault.description else {}), + **({"error_uri": fault.error_uri} if fault.error_uri else {}), + } + status_code: Final = 401 if fault.code == "invalid_client" else 400 + return JSONResponse(status_code=status_code, content=content, headers=TOKEN_NO_CACHE_HEADERS) + + def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: """RFC 6749 §5.2 response for a token-endpoint fault. Caller-actionable rejections relay the upstream's code on the status that code implies (401 for invalid_client per §5.2, else 400); @@ -42,13 +60,7 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: blamed for, or shown the internals of, a failure only the operator can fix.""" match fault.tag: case "caller_rejected": - content: Final = { - "error": fault.code, - **({"error_description": fault.description} if fault.description else {}), - **({"error_uri": fault.error_uri} if fault.error_uri else {}), - } - status_code = 401 if fault.code == "invalid_client" else 400 - return JSONResponse(status_code=status_code, content=content, headers=TOKEN_NO_CACHE_HEADERS) + return _render_caller_rejected(fault) case "gateway_rejected": return JSONResponse( status_code=502, @@ -65,6 +77,13 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: content={"error": fault.code, "error_description": description}, headers=TOKEN_NO_CACHE_HEADERS, ) + case "upstream_registration_refused": + return _render_caller_rejected( + CallerRejected( + code="unauthorized_client", + description=_registration_refused_description(fault.status_code), + ) + ) case "upstream_protocol_fault": return JSONResponse( status_code=502, @@ -78,7 +97,7 @@ def render_token_fault(fault: UpstreamOAuthFault) -> JSONResponse: def dcr_fault_detail(fault: UpstreamOAuthFault) -> tuple[int, str]: """Status and detail string for a registration fault, raised as HTTPException by the caller. RFC 7591 §3.2.2 defines registration errors as 400, so a contract-conformant rejection is 400 - regardless of the status the upstream chose; everything else is a 502 upstream fault.""" + regardless of the upstream status; a bare 401/403 is a registration refusal rendered as 403.""" match fault.tag: case "caller_rejected": detail: Final = f"{fault.code}: {fault.description}" if fault.description else fault.code @@ -87,6 +106,8 @@ def dcr_fault_detail(fault: UpstreamOAuthFault) -> tuple[int, str]: return 502, _gateway_rejected_description(fault.code) case "upstream_reported_fault": return _upstream_reported_status_and_description(fault.code) + case "upstream_registration_refused": + return 403, _registration_refused_description(fault.status_code) case "upstream_protocol_fault": return 502, fault.note case _: diff --git a/litellm/proxy/_experimental/mcp_server/faults/types.py b/litellm/proxy/_experimental/mcp_server/faults/types.py index 4b9505ad801..d081d9d735e 100644 --- a/litellm/proxy/_experimental/mcp_server/faults/types.py +++ b/litellm/proxy/_experimental/mcp_server/faults/types.py @@ -77,4 +77,12 @@ class UpstreamProtocolFault(BaseModel): note: str -UpstreamOAuthFault: TypeAlias = CallerRejected | GatewayRejected | UpstreamReportedFault | UpstreamProtocolFault +class UpstreamRegistrationRefused(BaseModel): + model_config = ConfigDict(frozen=True) + tag: Literal["upstream_registration_refused"] = "upstream_registration_refused" + status_code: Literal[401, 403] + + +UpstreamOAuthFault: TypeAlias = ( + CallerRejected | GatewayRejected | UpstreamReportedFault | UpstreamProtocolFault | UpstreamRegistrationRefused +) diff --git a/litellm/proxy/_experimental/mcp_server/mcp_debug.py b/litellm/proxy/_experimental/mcp_server/mcp_debug.py index ff30ef99ebd..1f157aefdc3 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_debug.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_debug.py @@ -100,14 +100,24 @@ Usage with curl:: http://localhost:4000/mcp/atlassian_mcp """ +import asyncio +import base64 +import io import json -from collections.abc import Callable, Mapping +import re +from collections.abc import AsyncIterator, Callable, Mapping +from http.cookies import CookieError, SimpleCookie +from itertools import islice from types import MappingProxyType from typing import Final +from urllib.parse import parse_qsl, quote, quote_plus, unquote_plus, urlencode +import httpx +from pydantic import JsonValue, TypeAdapter from starlette.requests import HTTPConnection from starlette.types import Message, Send +from litellm.litellm_core_utils.secret_redaction import REDACTED, redact_string from litellm.litellm_core_utils.sensitive_data_masker import SensitiveDataMasker from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution @@ -221,6 +231,10 @@ class MCPDebug: @staticmethod def _mask(value: str | None) -> str: """Mask a single value for safe display in headers.""" + return MCPDebug.mask_secret(value) + + @staticmethod + def mask_secret(value: str | None) -> str: if not value: return "(none)" return MCPDebug._masker._mask_value(value) @@ -378,3 +392,230 @@ class MCPDebug: server_url=server_url, server_auth_type=server_auth_type, ) + + +_BODY_PREVIEW_CHARS: Final = 512 +_BODY_CAPTURE_BYTES: Final = 16384 +_CAPTURE_TIMEOUT_SECONDS: Final = 1.0 +_CAPTURE_EXTENSION: Final = "litellm_mcp_error_preview" +_SAFE_HEADER_NAMES: Final = frozenset({"content-type", "content-length", "accept"}) +_PUBLIC_HEADER_NAMES: Final = _SAFE_HEADER_NAMES | frozenset(("host", "user-agent", "accept-encoding", "connection")) +_JSON_BODY: Final = TypeAdapter(JsonValue) +_LOG_MASKER: Final = SensitiveDataMasker(visible_prefix=0, visible_suffix=0) + + +def _safe_text(value: str, limit: int = _BODY_PREVIEW_CHARS) -> str: + escaped: Final = "".join(json.dumps(char)[1:-1] if ord(char) < 32 or ord(char) == 127 else char for char in value) + return escaped if len(escaped) <= limit else f"{escaped[:limit]}...(truncated)" + + +def safe_upstream_url(url: httpx.URL) -> str: + return _safe_text(str(url.copy_with(username="", password="", path="/", query=None, fragment=None))) + + +def _sensitive_field(key: str) -> bool: + normalized: Final = re.sub(r"[^a-z0-9]", "", key.casefold()) + return normalized in ("code", "clientassertion") or any( + pattern in normalized for pattern in _LOG_MASKER.sensitive_patterns + ) + + +def _redact_object( + fields: Mapping[str, JsonValue], +) -> dict[str, JsonValue]: # mutable-ok: the standard JSON encoder requires dict objects + return { # mutable-ok: construct the JSON object once for the standard parser and encoder + key: REDACTED if _sensitive_field(key) else value for key, value in fields.items() + } + + +def _header_secret_values(name: str, value: str) -> tuple[str, ...]: + if name == "cookie": + cookie: Final = SimpleCookie[str]() + try: + cookie.load(value) + except CookieError: + return (value,) + return (value, *(item.value for item in cookie.values())) + if name not in ("authorization", "proxy-authorization"): + return (value,) + scheme, _, credential = value.partition(" ") + if scheme.lower() != "basic": + return (value, credential) + try: + decoded: Final = base64.b64decode(credential, validate=True).decode("utf-8") + except ValueError: + return (value, credential) + password: Final = decoded.partition(":")[2] + return (value, credential, decoded, password, unquote_plus(password)) + + +def _body_secret_values(request: httpx.Request) -> tuple[str, ...] | None: + try: + raw: Final = request.content + except httpx.RequestNotRead: + return None + if not raw: + return () + if len(raw) > _BODY_CAPTURE_BYTES: + return None + if request.headers.get("content-type", "").split(";", 1)[0].strip().lower() == "application/x-www-form-urlencoded": + return tuple(value for key, value in parse_qsl(raw.decode("utf-8", errors="replace")) if _sensitive_field(key)) + try: + body: Final = _JSON_BODY.validate_json(raw) + except ValueError: + return None + from litellm.proxy._experimental.mcp_server.utils import ( # noqa: PLC0415 # MCP utils imports clients; inspect bodies only after initialization + json_string_leaves, + ) + + leaves: Final = json_string_leaves(body) + if leaves is None: + return None + return tuple( + value + for path, value in leaves + if not path or any(isinstance(part, str) and _sensitive_field(part) for part in path) + ) + + +def _request_secret_values(request: httpx.Request) -> tuple[str, ...] | None: + body_values: Final = _body_secret_values(request) + if body_values is None: + return None + values: Final = ( + *body_values, + request.url.password, + *(value for _, value in request.url.params.multi_items()), + *( + secret + for name, value in request.headers.items() + if name not in _PUBLIC_HEADER_NAMES + for secret in _header_secret_values(name, value) + ), + ) + return tuple(sorted(frozenset(value for value in values if value), key=len, reverse=True)) + + +def _mask_known_values(value: str, secrets: tuple[str, ...]) -> str: + variants: Final = tuple( + sorted( + frozenset( + variant + for secret in secrets + for variant in (secret, json.dumps(secret)[1:-1], quote(secret, safe=""), quote_plus(secret)) + ), + key=len, + reverse=True, + ) + ) + return re.sub("|".join(re.escape(secret) for secret in variants), REDACTED, value) if variants else value + + +def _preview(raw: bytes, content_type: str = "", secrets: tuple[str, ...] = ()) -> str: + if not raw: + return "(empty)" + if len(raw) > _BODY_CAPTURE_BYTES: + return "(omitted: body exceeds capture limit)" + try: + parsed: Final = _JSON_BODY.validate_python(json.loads(raw, object_hook=_redact_object)) + except (ValueError, RecursionError): + text: Final = raw.decode("utf-8", errors="replace") + if ( + content_type.split(";", 1)[0].strip().lower() != "application/x-www-form-urlencoded" + or "=" not in text + or any(char in text for char in "<>\n\r") + ): + return "(omitted: unstructured body)" + fields: Final = parse_qsl(text, keep_blank_values=True) + return _safe_text( + _mask_known_values( + urlencode(tuple((key, REDACTED if _sensitive_field(key) else value) for key, value in fields)), secrets + ) + ) + if not isinstance(parsed, (dict, list)): + return "(omitted: unstructured body)" + return _safe_text(redact_string(_mask_known_values(json.dumps(parsed, separators=(",", ":")), secrets))) + + +def _masked_headers(headers: httpx.Headers) -> str: + return _safe_text(", ".join(f"{name}={value}" for name, value in headers.items() if name in _SAFE_HEADER_NAMES)) + + +def _request_body_preview(request: httpx.Request, secrets: tuple[str, ...] | None) -> str: + try: + return _preview(request.content, request.headers.get("content-type", ""), secrets or ()) + except httpx.RequestNotRead: + return "(streamed, not captured)" + + +def _response_body_preview(response: httpx.Response, secrets: tuple[str, ...] | None) -> str: + if secrets is None: + return "(omitted: request credentials unavailable)" + captured: Final = response.extensions.get(_CAPTURE_EXTENSION) + if isinstance(captured, str): + return captured + try: + return _preview(response.content, response.headers.get("content-type", ""), secrets) + except httpx.ResponseNotRead: + return "(not read)" + + +async def _read_error_prefix(chunks: AsyncIterator[bytes], limit: int) -> bytes: + buffer: Final = io.BytesIO() + async for chunk in chunks: + buffer.write(chunk[: limit - buffer.tell()]) + if buffer.tell() >= limit: + break + return buffer.getvalue() + + +async def capture_upstream_error_response(response: httpx.Response) -> None: + if not response.is_error: + return + try: + prefix: Final = await asyncio.wait_for( + _read_error_prefix(response.aiter_bytes(chunk_size=4096), _BODY_CAPTURE_BYTES + 1), + timeout=_CAPTURE_TIMEOUT_SECONDS, + ) + response._content = prefix # pyright: ignore[reportPrivateUsage] # rebind-ok: httpx has no public setter to retain consumed bytes for auth retries + secrets: Final = _request_secret_values(response.request) + preview: Final = ( + _preview(prefix, response.headers.get("content-type", ""), secrets) + if secrets is not None + else "(omitted: request credentials unavailable)" + ) + except (asyncio.TimeoutError, httpx.HTTPError, httpx.StreamError): + response._content = b"" # pyright: ignore[reportPrivateUsage] # rebind-ok: httpx auth retries must survive diagnostic read failures + response.extensions[_CAPTURE_EXTENSION] = ( + "(unavailable: error body read failed)" # rebind-ok: httpx response hooks communicate through extensions + ) + return + response.extensions[_CAPTURE_EXTENSION] = preview # rebind-ok: httpx response hooks communicate through extensions + + +def describe_upstream_response(response: httpx.Response) -> str: + try: + request: Final = response.request + except RuntimeError: + return f"HTTP {response.status_code} | request unavailable" + secrets: Final = _request_secret_values(request) + return ( + f"{_safe_text(request.method)} {safe_upstream_url(request.url)} -> HTTP {response.status_code}" + f" | request headers: {_masked_headers(request.headers)}" + f" | request body: {_request_body_preview(request, secrets)}" + f" | response body: {_response_body_preview(response, secrets)}" + ) + + +def describe_upstream_http_failure(exc: BaseException) -> str | None: + from litellm.proxy._experimental.mcp_server.faults.traversal import ( # noqa: PLC0415 # fault package initialization imports the credential resolver + iter_exception_tree, + ) + + lines: Final = tuple( + describe_upstream_response(response) + for current in islice(iter_exception_tree(exc), 16) + for response in (getattr(current, "response", None),) + if isinstance(response, httpx.Response) + ) + return " | ".join(lines) or None diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 17d3098bf4c..cc7ea0c1ea5 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -81,7 +81,7 @@ from litellm.proxy._experimental.mcp_server.faults.list_outcomes import ( raise_classified_list_failure, upstream_auth_challenge, ) -from litellm.proxy._experimental.mcp_server.mcp_debug import record_auth_resolution +from litellm.proxy._experimental.mcp_server.mcp_debug import describe_upstream_http_failure, record_auth_resolution from litellm.proxy._experimental.mcp_server.oauth2_token_cache import ( MCPPerUserTokenCache, mcp_per_user_token_cache, @@ -255,6 +255,7 @@ _TRUE_ENV_VALUES: Final = frozenset(("1", "true", "yes", "on")) _OAUTH_DISCOVERY_RETRY_DELAYS_SECONDS: Final = (0.05, 0.15) _OAUTH_DISCOVERY_RETRY_BASE_SECONDS: Final = 30.0 _OAUTH_DISCOVERY_RETRY_MAX_SECONDS: Final = 900.0 +_OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS: Final = 300.0 def _oauth_discovery_now() -> float: @@ -1404,6 +1405,11 @@ def _extract_upstream_auth_failure( return upstream_auth_challenge(exc) +def _upstream_failure_suffix(exc: BaseException) -> str: + detail: Final = describe_upstream_http_failure(exc) + return f"\n upstream exchange: {detail}" if detail else "" + + def _obo_retry_applies(server: MCPServer, subject_token: str | None) -> bool: """Whether an upstream 401/403 should invalidate the minted credential and retry once. @@ -1947,6 +1953,10 @@ class MCPServerManager: slot: Final = self._oauth_discovery_slot(server_id) return slot is not None and slot.generation == generation + def _expire_temporary_oauth_discovery(self, server_id: str, generation: int) -> None: + if self._oauth_discovery_slot_is_current(server_id, generation): + self._remove_oauth_discovery_slot(server_id) + def _publish_resolved_oauth_server( self, server: MCPServer, @@ -1959,7 +1969,13 @@ class MCPServerManager: elif server.server_id in self.config_mcp_servers: self.config_mcp_servers[server.server_id] = server else: - return None + asyncio.get_running_loop().call_later( + _OAUTH_TEMPORARY_DISCOVERY_TTL_SECONDS, + self._expire_temporary_oauth_discovery, + server.server_id, + generation, + ) + return server self._remove_oauth_discovery_slot(server.server_id) return server @@ -2051,6 +2067,12 @@ class MCPServerManager: if slot.task is not None: if not slot.task.done() or _oauth_discovery_now() < slot.retry_not_before: return slot.task, slot.generation + if ( + not slot.task.cancelled() + and slot.task.exception() is None + and isinstance(slot.task.result(), _OAuthDiscoveryResolved) + ): + return slot.task, slot.generation task: Final = asyncio.create_task( self._run_oauth_metadata_resolution(self._registered_server(server), slot.generation) ) @@ -2080,7 +2102,7 @@ class MCPServerManager: if should_defer != has_slot: self._set_oauth_discovery_deferred(server.server_id, should_defer) - async def ensure_oauth_metadata_discovered(self, server: MCPServer) -> MCPServer: + async def ensure_oauth_metadata_discovered(self, server: MCPServer, *, _retry_stale: bool = True) -> MCPServer: """Join the bounded discovery task and return the resolved server. Concurrent callers share one task per server. A failed attempt remains @@ -2107,13 +2129,13 @@ class MCPServerManager: outcome: Final = await asyncio.shield(task) except asyncio.CancelledError: if task.cancelled() and not self._oauth_discovery_slot_is_current(server.server_id, generation): - return await self.ensure_oauth_metadata_discovered(server) + return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale) raise match outcome: case _OAuthDiscoveryResolved(resolved_server): return resolved_server case _OAuthDiscoveryStale(): - return await self.ensure_oauth_metadata_discovered(server) + return await self._rejoin_oauth_metadata_discovery(server, retry_stale=_retry_stale) case _OAuthDiscoveryFailed(timed_out=timed_out): current: Final = self._registered_server(server) if current.is_client_forwarded_token: @@ -2125,6 +2147,14 @@ class MCPServerManager: detail=f"OAuth metadata discovery {reason} for MCP server {server_ref!r}", ) + async def _rejoin_oauth_metadata_discovery(self, server: MCPServer, *, retry_stale: bool) -> MCPServer: + if retry_stale: + return await self.ensure_oauth_metadata_discovered(server, _retry_stale=False) + current: Final = self._registered_server(server) + if not _oauth_endpoints_unresolved(current) or current.is_client_forwarded_token: + return current + raise HTTPException(status_code=503, detail="OAuth metadata discovery changed repeatedly; retry shortly") + def _remember_upstream_initialize_instructions(self, server: MCPServer, client: MCPClient) -> None: raw: Final[str | None] = getattr(client, "_last_initialize_instructions", None) if raw and str(raw).strip(): @@ -4340,7 +4370,9 @@ class MCPServerManager: except MCPServerListError: raise except Exception as e: - verbose_logger.warning("Failed to get tools from server %s: %s", server.name, e) + verbose_logger.warning( + "Failed to get tools from server %s: %s%s", server.name, type(e).__name__, _upstream_failure_suffix(e) + ) raise_classified_list_failure(e, server.name, suppress_challenge=server.is_dcr_bridge) async def get_prompts_from_server( @@ -5085,7 +5117,9 @@ class MCPServerManager: verbose_logger.warning("Connection error while listing tools from %s: %s", server_name, e) raise MCPServerListError(ServerListFault(tag="unreachable"), server_name) from e except Exception as e: - verbose_logger.warning("Error listing tools from %s: %s", server_name, e) + verbose_logger.warning( + "Error listing tools from %s: %s%s", server_name, type(e).__name__, _upstream_failure_suffix(e) + ) raise_classified_list_failure(e, server_name) _SHORT_PREFIX_MAX_REHASH_ATTEMPTS = 1024 diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py index ad18d1bb10f..43d97abe4db 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py @@ -37,6 +37,7 @@ import httpx from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError from typing_extensions import assert_never +from litellm._logging import verbose_logger from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( InMemoryTokenCacheBackend, OAuthToken, @@ -101,6 +102,11 @@ async def post_client_credentials_grant( from litellm.llms.custom_httpx.http_handler import ( # noqa: PLC0415 # defer heavy handler import to call time get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # handler factory params are coarsely typed ) + from litellm.proxy._experimental.mcp_server.mcp_debug import ( # noqa: PLC0415 # diagnostics import credential enums through this package + describe_upstream_http_failure, + describe_upstream_response, + safe_upstream_url, + ) from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 # deferred with the handler import try: @@ -110,15 +116,28 @@ async def post_client_credentials_grant( ) except httpx.HTTPStatusError as status_err: status_code: Final = status_err.response.status_code + verbose_logger.warning( + "OAuth2 client_credentials token request denied:\n upstream exchange: %s", + describe_upstream_http_failure(status_err), + ) return TokenEndpointDenied(status_code=status_code, detail=f"token endpoint returned HTTP {status_code}") except Exception as exc: # noqa: BLE001 # any transport failure is the same outcome: unreachable - return TokenEndpointUnreachable(detail=str(exc)) + verbose_logger.warning( + "OAuth2 client_credentials POST %s failed: %s", safe_upstream_url(httpx.URL(url)), type(exc).__name__ + ) + return TokenEndpointUnreachable(detail=type(exc).__name__) try: body: Final = _TOKEN_BODY_ADAPTER.validate_json(response.content) except ValidationError: + verbose_logger.warning("OAuth2 client_credentials invalid response: %s", describe_upstream_response(response)) return TokenEndpointDenied( status_code=response.status_code, detail="token endpoint returned a non-JSON-object body" ) + access_token: Final = body.get("access_token") + if not isinstance(access_token, str) or not access_token: + verbose_logger.warning( + "OAuth2 client_credentials response has no access token | %s", describe_upstream_response(response) + ) return TokenEndpointSuccess(body=body) diff --git a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py index 27cb632c843..7a97e995570 100644 --- a/litellm/proxy/_experimental/mcp_server/rest_endpoints.py +++ b/litellm/proxy/_experimental/mcp_server/rest_endpoints.py @@ -191,6 +191,7 @@ if MCP_AVAILABLE: execute_mcp_tool, filter_tools_by_allowed_tools, filter_tools_by_key_team_permissions, + fire_mcp_tool_call_failure_logging, ) ######################################################## @@ -232,6 +233,20 @@ if MCP_AVAILABLE: return result return outcome + async def _safe_fire_mcp_tool_call_failure_logging( + logging_obj: "LiteLLMLoggingObj | None", + exception: Exception, + start_time: datetime, + user_api_key_auth: UserAPIKeyAuth, + request_data: Mapping[str, object], + ) -> None: + try: + await fire_mcp_tool_call_failure_logging( + logging_obj, exception, start_time, user_api_key_auth, request_data + ) + except Exception as logging_error: + verbose_logger.warning("MCP tool call failure logging failed (continuing): %s", logging_error) + def _relay_upstream_auth_http_exception(e: MCPUpstreamAuthError, request: Request) -> HTTPException: """Convert a client-forwarded pass-through upstream 401 into an HTTPException that preserves the upstream WWW-Authenticate, so a standards-compliant MCP client can run the upstream OAuth flow @@ -310,26 +325,39 @@ if MCP_AVAILABLE: ) # MCP_TOOL_CALL_TOOL_NAME: run the same pre-call pipeline as the normal path so the tool # execution is spend-logged and guardrail-checked. - (_, virtual_logging_obj) = await ProxyBaseLLMRequestProcessing(data=data).common_processing_pre_call_logic( - request=request, - user_api_key_dict=user_api_key_dict, - proxy_config=proxy_config, - route_type=CallTypes.call_mcp_tool.value, - proxy_logging_obj=proxy_logging_obj, - general_settings=general_settings, - ) - _tool_start_time: Final = datetime.now() - result: Final = await handle_mcp_tool_call( - tool_name=tool_arguments.get("tool_name", ""), - arguments=tool_arguments.get("arguments") or {}, - user_api_key_dict=user_api_key_dict, - client_ip=rest_client_ip, - mcp_auth_header=virtual_mcp_auth_header, - mcp_server_auth_headers=virtual_mcp_server_auth_headers, - oauth2_headers=virtual_oauth2_headers, - raw_headers=virtual_raw_headers, - litellm_logging_obj=virtual_logging_obj, - ) + virtual_processor: Final = ProxyBaseLLMRequestProcessing(data=data) + _request_start_time: Final = datetime.now() # noqa: DTZ005 # naive to match the tool start time below + try: + (_, virtual_logging_obj) = await virtual_processor.common_processing_pre_call_logic( + request=request, + user_api_key_dict=user_api_key_dict, + proxy_config=proxy_config, + route_type=CallTypes.call_mcp_tool.value, + proxy_logging_obj=proxy_logging_obj, + general_settings=general_settings, + ) + _tool_start_time: Final = datetime.now() + result: Final = await handle_mcp_tool_call( + tool_name=tool_arguments.get("tool_name", ""), + arguments=tool_arguments.get("arguments") or {}, + user_api_key_dict=user_api_key_dict, + client_ip=rest_client_ip, + mcp_auth_header=virtual_mcp_auth_header, + mcp_server_auth_headers=virtual_mcp_server_auth_headers, + oauth2_headers=virtual_oauth2_headers, + raw_headers=virtual_raw_headers, + litellm_logging_obj=virtual_logging_obj, + ) + except Exception as e: + virtual_request_data: Final = virtual_processor.data + await _safe_fire_mcp_tool_call_failure_logging( + virtual_request_data.get("litellm_logging_obj"), + e, + _request_start_time, + user_api_key_dict, + virtual_request_data, + ) + raise return await _safe_fire_mcp_tool_call_logging( virtual_logging_obj, result, @@ -965,7 +993,9 @@ if MCP_AVAILABLE: apply_tool_filters=apply_tool_filters, ) except Exception as e: - verbose_logger.exception("Error getting tools from %s: %s", server.name, e) + verbose_logger.warning( + "Error getting tools from %s: %s", server.name, classify_list_exception(e).tag + ) return (), classify_list_exception(e) return tools_result, ServerListOk(tool_count=len(tools_result)) @@ -1079,65 +1109,73 @@ if MCP_AVAILABLE: ) proxy_base_llm_response_processor: Final = ProxyBaseLLMRequestProcessing(data=data) - ( - data, - logging_obj, - ) = await proxy_base_llm_response_processor.common_processing_pre_call_logic( - request=request, - user_api_key_dict=user_api_key_dict, - proxy_config=proxy_config, - route_type=CallTypes.call_mcp_tool.value, - proxy_logging_obj=proxy_logging_obj, - general_settings=general_settings, - ) + _request_start_time: Final = datetime.now() # noqa: DTZ005 # naive to match the tool start time below + try: + ( + data, + logging_obj, + ) = await proxy_base_llm_response_processor.common_processing_pre_call_logic( + request=request, + user_api_key_dict=user_api_key_dict, + proxy_config=proxy_config, + route_type=CallTypes.call_mcp_tool.value, + proxy_logging_obj=proxy_logging_obj, + general_settings=general_settings, + ) - # Extract MCP auth headers from request and add to data dict - ( - mcp_auth_header, - mcp_server_auth_headers, - raw_headers_from_request, - ) = _extract_mcp_headers_from_request(request, MCPRequestHandler) - if mcp_auth_header: - data["mcp_auth_header"] = mcp_auth_header - if mcp_server_auth_headers: - data["mcp_server_auth_headers"] = mcp_server_auth_headers - data["raw_headers"] = raw_headers_from_request + # Extract MCP auth headers from request and add to data dict + ( + mcp_auth_header, + mcp_server_auth_headers, + raw_headers_from_request, + ) = _extract_mcp_headers_from_request(request, MCPRequestHandler) + if mcp_auth_header: + data["mcp_auth_header"] = mcp_auth_header + if mcp_server_auth_headers: + data["mcp_server_auth_headers"] = mcp_server_auth_headers + data["raw_headers"] = raw_headers_from_request - # Extract user_api_key_auth from metadata and add to top level - # call_mcp_tool expects user_api_key_auth as a top-level parameter - if "metadata" in data and "user_api_key_auth" in data["metadata"]: - data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"] + # Extract user_api_key_auth from metadata and add to top level + # call_mcp_tool expects user_api_key_auth as a top-level parameter + if "metadata" in data and "user_api_key_auth" in data["metadata"]: + data["user_api_key_auth"] = data["metadata"]["user_api_key_auth"] - # Resolve allowed MCP servers with IP filtering - ( - allowed_mcp_servers, - canonical_server_id, - ) = await _resolve_allowed_mcp_servers_with_ip_filter(request, user_api_key_dict, server_id) + # Resolve allowed MCP servers with IP filtering + ( + allowed_mcp_servers, + canonical_server_id, + ) = await _resolve_allowed_mcp_servers_with_ip_filter(request, user_api_key_dict, server_id) - # Look up per-user OAuth headers for this server (mirrors list_tool_rest_api). - user_oauth_extra_headers: dict[str, str] | None = None - target_server: Final = next( - (s for s in allowed_mcp_servers if s.server_id == canonical_server_id), - None, - ) - if target_server is not None: - user_oauth_extra_headers = await _get_user_oauth_extra_headers(target_server, user_api_key_dict) + # Look up per-user OAuth headers for this server (mirrors list_tool_rest_api). + user_oauth_extra_headers: dict[str, str] | None = None + target_server: Final = next( + (s for s in allowed_mcp_servers if s.server_id == canonical_server_id), + None, + ) + if target_server is not None: + user_oauth_extra_headers = await _get_user_oauth_extra_headers(target_server, user_api_key_dict) - # Call execute_mcp_tool directly (permission checks already done) - _tool_start_time: Final = datetime.now() - result: Final = await execute_mcp_tool( - name=tool_name, - arguments=tool_arguments, - allowed_mcp_servers=allowed_mcp_servers, - start_time=_tool_start_time, - user_api_key_auth=data.get("user_api_key_auth"), - mcp_auth_header=data.get("mcp_auth_header"), - mcp_server_auth_headers=data.get("mcp_server_auth_headers"), - oauth2_headers=user_oauth_extra_headers or data.get("oauth2_headers"), - raw_headers=data.get("raw_headers"), - litellm_logging_obj=data.get("litellm_logging_obj"), - requested_server_id=canonical_server_id, - ) + # Call execute_mcp_tool directly (permission checks already done) + _tool_start_time: Final = datetime.now() + result: Final = await execute_mcp_tool( + name=tool_name, + arguments=tool_arguments, + allowed_mcp_servers=allowed_mcp_servers, + start_time=_tool_start_time, + user_api_key_auth=data.get("user_api_key_auth"), + mcp_auth_header=data.get("mcp_auth_header"), + mcp_server_auth_headers=data.get("mcp_server_auth_headers"), + oauth2_headers=user_oauth_extra_headers or data.get("oauth2_headers"), + raw_headers=data.get("raw_headers"), + litellm_logging_obj=data.get("litellm_logging_obj"), + requested_server_id=canonical_server_id, + ) + except Exception as e: + request_data: Final = proxy_base_llm_response_processor.data + await _safe_fire_mcp_tool_call_failure_logging( + request_data.get("litellm_logging_obj"), e, _request_start_time, user_api_key_dict, request_data + ) + raise return await _safe_fire_mcp_tool_call_logging( logging_obj, result, diff --git a/litellm/proxy/_experimental/mcp_server/server.py b/litellm/proxy/_experimental/mcp_server/server.py index 95129d9aeed..43a7f894db8 100644 --- a/litellm/proxy/_experimental/mcp_server/server.py +++ b/litellm/proxy/_experimental/mcp_server/server.py @@ -3339,6 +3339,43 @@ if MCP_AVAILABLE: ) return result + async def fire_mcp_tool_call_failure_logging( + logging_obj: LiteLLMLoggingObj | None, + exception: Exception, + start_time: datetime, + user_api_key_auth: UserAPIKeyAuth | None, + request_data: Mapping[str, object], + ) -> None: + """Failure logging shared by the ``/mcp`` path and the REST endpoint. Call from + inside the ``except`` block so the traceback is still available. + + The failure handlers run first because ``_ProxyDBLogger.async_post_call_failure_hook`` + builds the failure spend-log row from the ``standard_logging_object`` they produce; + both gate on ``should_run_logging``, so the ``@client`` wrapper does not log twice. + A relayed upstream 401 (``MCPUpstreamAuthError``) is an expected caller-must-reauth + signal and skips ``post_call_failure_hook``, which fires the ``llm_exceptions`` alert. + """ + from litellm.proxy.proxy_server import proxy_logging_obj + + traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) + if logging_obj is not None: + end_time: Final = datetime.now() # noqa: DTZ005 # naive to match `start_time`, which it is subtracted from + logging_obj.failure_handler(exception, traceback_str, start_time, end_time) + await logging_obj.async_failure_handler(exception, traceback_str, start_time, end_time) + + if isinstance(exception, MCPUpstreamAuthError) or not proxy_logging_obj or user_api_key_auth is None: + return + sanitized_request_data: Final = { + key: value for key, value in request_data.items() if key not in _MCP_CREDENTIAL_REQUEST_FIELDS + } + await proxy_logging_obj.post_call_failure_hook( + request_data=sanitized_request_data, + original_exception=exception, + user_api_key_dict=user_api_key_auth, + route="/mcp/call_tool", + traceback_str=traceback_str, + ) + @client async def call_mcp_tool( name: str, @@ -3405,40 +3442,8 @@ if MCP_AVAILABLE: raw_headers=raw_headers, **kwargs, ) - except MCPUpstreamAuthError: - # A client-forwarded pass-through upstream 401 is an expected caller-must-reauth signal, so - # re-raise it without post_call_failure_hook, which fires the proxy's llm_exceptions alert. - # mcp_server_tool_call then downgrades it to an informational isError result for the - # streamable client. Note: this function is @client-decorated, so the decorator's standard - # failure logging still records the event (spend log / OTel); only the extra alert sink is - # skipped here. - raise except Exception as e: - traceback_str: Final = traceback.format_exc(limit=MAXIMUM_TRACEBACK_LINES_TO_LOG) - from litellm.proxy.proxy_server import proxy_logging_obj - - # Ordering is load-bearing. ``_ProxyDBLogger.async_post_call_failure_hook``, - # reached below, writes the failure spend-log row from this logger's - # ``standard_logging_object``, which only exists once the failure handlers - # have run. Flush them first or the row lands with - # ``guardrail_information=None`` and a guardrail block is never counted. - # - # Not double-logged: both handlers gate on ``should_run_logging`` and then - # mark it, so the ``@client`` wrapper's own post-raise logging no-ops on this - # logger, same as ``_fire_mcp_tool_call_logging`` does for ``isError=True``. - if litellm_logging_obj is not None: - end_time: Final = datetime.now() # noqa: DTZ005 # naive to match `start_time`, which it is subtracted from - litellm_logging_obj.failure_handler(e, traceback_str, start_time, end_time) - await litellm_logging_obj.async_failure_handler(e, traceback_str, start_time, end_time) - - if proxy_logging_obj and user_api_key_auth: - await proxy_logging_obj.post_call_failure_hook( - request_data=kwargs, - original_exception=e, - user_api_key_dict=user_api_key_auth, - route="/mcp/call_tool", - traceback_str=traceback_str, - ) + await fire_mcp_tool_call_failure_logging(litellm_logging_obj, e, start_time, user_api_key_auth, kwargs) raise if litellm_logging_obj: diff --git a/litellm/proxy/_lazy_features.py b/litellm/proxy/_lazy_features.py index 50e0a961a49..dd1180b30ad 100644 --- a/litellm/proxy/_lazy_features.py +++ b/litellm/proxy/_lazy_features.py @@ -8,14 +8,17 @@ omits each feature's routes until the feature is warmed. import asyncio import importlib -from collections.abc import Callable +from collections.abc import Callable, Mapping, Sequence from collections.abc import Set as AbstractSet from dataclasses import dataclass, field +from types import MappingProxyType from typing import TYPE_CHECKING, Final +from starlette.routing import BaseRoute, Match from starlette.types import Receive, Scope, Send from litellm._logging import verbose_proxy_logger +from litellm.proxy.route_priority import hot_routes_first if TYPE_CHECKING: from fastapi import APIRouter, FastAPI @@ -185,6 +188,31 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = ( module_path="litellm.proxy.management_endpoints.config_override_endpoints", path_prefixes=("/config_overrides",), ), + LazyFeature( + name="llm_passthrough", + module_path="litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints", + path_prefixes=( + "/anthropic/", + "/assemblyai/", + "/azure/", + "/azure_ai/", + "/bedrock/", + "/cohere/", + "/comprehendmedical", + "/cursor/", + "/eu.assemblyai/", + "/gemini/", + "/gigachat/", + "/milvus/", + "/mistral/", + "/openai/", + "/openai_passthrough/", + "/vertex-ai/", + "/vertex_ai/", + "/vllm/", + "/watsonx/", + ), + ), LazyFeature( name="realtime", module_path="litellm.proxy.realtime_endpoints.endpoints", @@ -308,14 +336,73 @@ class LazyFeatureMiddleware: if root_path and path.startswith(root_path + "/"): path = path[len(root_path) :] # rebind-ok: local strip after the boundary check above for feat in self._features: - if feat.module_path in self._loaded: + if feat.module_path in self._loaded or not feat.matches(path): continue - if feat.matches(path): - await _force_load(self._fastapi_app, feat) + if _eager_route_wins(self._fastapi_app, feat, scope): + continue + await _force_load(self._fastapi_app, feat, self._features) await self.app(scope, receive, send) -async def _force_load(app: "FastAPI", feat: LazyFeature) -> bool: +def _lazy_slots(app: "FastAPI") -> Mapping[str, BaseRoute | None]: + return app.state.lazy_slots if hasattr(app.state, "lazy_slots") else MappingProxyType({}) + + +def reserve_lazy_slot(app: "FastAPI", name: str, features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> None: + """Record the route the feature's router used to be included after, so its routes + are spliced back in there once it loads and keep the same precedence. Anchoring on + the route rather than its index survives later reordering of the table.""" + feat: Final = next(f for f in features if f.name == name) + anchor: Final = app.router.routes[-1] if app.router.routes else None + app.state.lazy_slots = MappingProxyType({**_lazy_slots(app), feat.module_path: anchor}) + + +def _slot_index(routes: Sequence[BaseRoute], anchor: BaseRoute | None) -> int: + if anchor is None: + return 0 + return next((i + 1 for i, route in enumerate(routes) if route is anchor), len(routes)) + + +def _eager_route_wins(app: "FastAPI", feat: LazyFeature, scope: Scope) -> bool: + """Routes ahead of a feature's reserved slot beat its routes in Starlette's scan, + so a request one of them fully matches never needs the feature loaded.""" + slots: Final = _lazy_slots(app) + if feat.module_path not in slots: + return False + ahead: Final = app.router.routes[: _slot_index(app.router.routes, slots[feat.module_path])] + return any(route.matches(scope)[0] is Match.FULL for route in ahead) + + +def _in_registry_order( + routes: Sequence[BaseRoute], + lazy_routes: Mapping[str, tuple[BaseRoute, ...]], + features: tuple[LazyFeature, ...], + slots: Mapping[str, BaseRoute | None], +) -> tuple[BaseRoute, ...]: + """Lazy routers land in registry order, not first-request order, so overlapping + paths (/openai/{endpoint:path} vs /openai/v1/realtime/calls) resolve the same + way no matter which feature a deployment happens to hit first. Features with a + reserved slot go back where they were eagerly included; the rest follow every + eager route.""" + rank: Final = MappingProxyType({f.module_path: i for i, f in enumerate(features)}) + modules: Final = tuple(sorted(lazy_routes, key=lambda m: rank.get(m, len(rank)))) + lazy_ids: Final = frozenset(id(route) for module_path in modules for route in lazy_routes[module_path]) + eager: Final = tuple(route for route in routes if id(route) not in lazy_ids) + + def slot_of(module_path: str) -> int: + return _slot_index(eager, slots[module_path]) if module_path in slots else len(eager) + + return tuple( + route + for index in range(len(eager) + 1) + for route in ( + *(r for module_path in modules if slot_of(module_path) == index for r in lazy_routes[module_path]), + *eager[index : index + 1], + ) + ) + + +async def _force_load(app: "FastAPI", feat: LazyFeature, features: tuple[LazyFeature, ...] = LAZY_FEATURES) -> bool: """Import + register a lazy feature exactly once per (app, module). Shared by the middleware and the /lazy/warm endpoint.""" if not hasattr(app.state, "lazy_loaded"): @@ -330,7 +417,18 @@ async def _force_load(app: "FastAPI", feat: LazyFeature) -> bool: # mutates app.router.routes, so it stays on the loop thread. loop: Final = asyncio.get_running_loop() module: Final = await loop.run_in_executor(None, importlib.import_module, feat.module_path) + before: Final = len(app.router.routes) feat.register_fn(app, module) + previous: Final[Mapping[str, tuple[BaseRoute, ...]]] = ( + app.state.lazy_routes if hasattr(app.state, "lazy_routes") else MappingProxyType({}) + ) + lazy_routes: Final[Mapping[str, tuple[BaseRoute, ...]]] = MappingProxyType( + {**previous, feat.module_path: tuple(app.router.routes[before:])} + ) + app.state.lazy_routes = lazy_routes # rebind-ok: the app owns the record of which routes each feature added + app.router.routes[:] = hot_routes_first( # rebind-ok: the app owns its route table + _in_registry_order(app.router.routes, lazy_routes, features, _lazy_slots(app)) + ) app.state.lazy_loaded.add(feat.module_path) app.openapi_schema = None verbose_proxy_logger.info( diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index f0af17ab818..53af85baac6 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -4335,7 +4335,7 @@ "/anthropic/{endpoint}": { "delete": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete_2", "parameters": [ { "in": "path", @@ -4379,7 +4379,7 @@ }, "get": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__get", + "operationId": "anthropic_proxy_route_anthropic__endpoint__get_2", "parameters": [ { "in": "path", @@ -4423,7 +4423,7 @@ }, "patch": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__patch", + "operationId": "anthropic_proxy_route_anthropic__endpoint__patch_2", "parameters": [ { "in": "path", @@ -4467,7 +4467,7 @@ }, "post": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__post", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post_2", "parameters": [ { "in": "path", @@ -4511,7 +4511,7 @@ }, "put": { "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", - "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "operationId": "anthropic_proxy_route_anthropic__endpoint__put_2", "parameters": [ { "in": "path", @@ -15836,6 +15836,4976 @@ } } }, + "llm_passthrough": { + "components": { + "schemas": { + "Body_image_edit_api_openai_deployments__model__images_edits_post": { + "properties": { + "image": { + "anyOf": [ + { + "items": { + "contentMediaType": "application/octet-stream", + "type": "string" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Image" + }, + "image[]": { + "anyOf": [ + { + "items": { + "contentMediaType": "application/octet-stream", + "type": "string" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Image[]" + }, + "mask": { + "anyOf": [ + { + "items": { + "contentMediaType": "application/octet-stream", + "type": "string" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Mask" + }, + "mask[]": { + "anyOf": [ + { + "items": { + "contentMediaType": "application/octet-stream", + "type": "string" + }, + "type": "array" + }, + { + "type": "null" + } + ], + "title": "Mask[]" + } + }, + "title": "Body_image_edit_api_openai_deployments__model__images_edits_post", + "type": "object" + }, + "ErrorResponse": { + "properties": { + "detail": { + "additionalProperties": true, + "example": { + "error": { + "code": "error_code", + "message": "Error message", + "param": "error_param", + "type": "error_type" + } + }, + "title": "Detail", + "type": "object" + } + }, + "required": [ + "detail" + ], + "title": "ErrorResponse", + "type": "object" + }, + "HTTPValidationError": { + "properties": { + "detail": { + "items": { + "$ref": "#/components/schemas/ValidationError" + }, + "title": "Detail", + "type": "array" + } + }, + "title": "HTTPValidationError", + "type": "object" + }, + "RealtimeClientSecretResponse": { + "description": "Response from POST /v1/realtime/client_secrets.\n\nBoth the top-level `value` and `session.client_secret.value`\nwill contain the encrypted token instead of the raw ephemeral key.\nThe `session` field is kept as a raw dict so unknown fields pass through.", + "properties": { + "expires_at": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Expires At" + }, + "session": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Session" + }, + "value": { + "title": "Value", + "type": "string" + } + }, + "required": [ + "value" + ], + "title": "RealtimeClientSecretResponse", + "type": "object" + }, + "RealtimeTranscriptionSessionResponse": { + "additionalProperties": true, + "description": "Response from POST /v1/realtime/transcription_sessions.\n\n`client_secret.value` contains the encrypted token instead of the raw\nephemeral key. Unknown fields pass through unchanged.", + "properties": { + "client_secret": { + "anyOf": [ + { + "additionalProperties": true, + "type": "object" + }, + { + "type": "null" + } + ], + "title": "Client Secret" + } + }, + "title": "RealtimeTranscriptionSessionResponse", + "type": "object" + }, + "ValidationError": { + "properties": { + "ctx": { + "title": "Context", + "type": "object" + }, + "input": { + "title": "Input" + }, + "loc": { + "items": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "integer" + } + ] + }, + "title": "Location", + "type": "array" + }, + "msg": { + "title": "Message", + "type": "string" + }, + "type": { + "title": "Error Type", + "type": "string" + } + }, + "required": [ + "loc", + "msg", + "type" + ], + "title": "ValidationError", + "type": "object" + } + } + }, + "paths": { + "/anthropic/{endpoint}": { + "delete": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", + "operationId": "anthropic_proxy_route_anthropic__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Anthropic Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", + "operationId": "anthropic_proxy_route_anthropic__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Anthropic Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", + "operationId": "anthropic_proxy_route_anthropic__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Anthropic Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", + "operationId": "anthropic_proxy_route_anthropic__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Anthropic Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/anthropic_completion)", + "operationId": "anthropic_proxy_route_anthropic__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Anthropic Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/assemblyai/{endpoint}": { + "delete": { + "operationId": "assemblyai_proxy_route_assemblyai__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "operationId": "assemblyai_proxy_route_assemblyai__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "operationId": "assemblyai_proxy_route_assemblyai__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "operationId": "assemblyai_proxy_route_assemblyai__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "operationId": "assemblyai_proxy_route_assemblyai__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/azure/{endpoint}": { + "delete": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/azure_ai/{endpoint}": { + "delete": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure_ai__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure_ai__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure_ai__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure_ai__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Call any azure endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/azure/{endpoint:path}`\n\nChecks if the deployment id in the url is a litellm model name. If so, it will route using the llm_router.allm_passthrough_route.", + "operationId": "azure_proxy_route_azure_ai__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Azure Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/bedrock/{endpoint}": { + "delete": { + "description": "This is the v1 passthrough for Bedrock.\nV2 is handled by the `/bedrock/v2` endpoint.\n[Docs](https://docs.litellm.ai/docs/pass_through/bedrock)", + "operationId": "bedrock_proxy_route_bedrock__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Bedrock Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "This is the v1 passthrough for Bedrock.\nV2 is handled by the `/bedrock/v2` endpoint.\n[Docs](https://docs.litellm.ai/docs/pass_through/bedrock)", + "operationId": "bedrock_proxy_route_bedrock__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Bedrock Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "This is the v1 passthrough for Bedrock.\nV2 is handled by the `/bedrock/v2` endpoint.\n[Docs](https://docs.litellm.ai/docs/pass_through/bedrock)", + "operationId": "bedrock_proxy_route_bedrock__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Bedrock Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "This is the v1 passthrough for Bedrock.\nV2 is handled by the `/bedrock/v2` endpoint.\n[Docs](https://docs.litellm.ai/docs/pass_through/bedrock)", + "operationId": "bedrock_proxy_route_bedrock__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Bedrock Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "This is the v1 passthrough for Bedrock.\nV2 is handled by the `/bedrock/v2` endpoint.\n[Docs](https://docs.litellm.ai/docs/pass_through/bedrock)", + "operationId": "bedrock_proxy_route_bedrock__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Bedrock Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/cohere/{endpoint}": { + "delete": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/cohere)", + "operationId": "cohere_proxy_route_cohere__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cohere Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/cohere)", + "operationId": "cohere_proxy_route_cohere__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cohere Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/cohere)", + "operationId": "cohere_proxy_route_cohere__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cohere Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/cohere)", + "operationId": "cohere_proxy_route_cohere__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cohere Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/cohere)", + "operationId": "cohere_proxy_route_cohere__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cohere Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/comprehendmedical": { + "post": { + "description": "AWS-SDK-shaped pass-through for Amazon Comprehend Medical: point the SDK's\n`endpoint_url` at `/comprehendmedical` and the operation is read from the\n`X-Amz-Target` header, per the AWS JSON 1.1 protocol.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical)", + "operationId": "comprehend_medical_sdk_proxy_route_comprehendmedical_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Comprehend Medical Sdk Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/comprehendmedical/{operation}": { + "post": { + "description": "Pass-through for Amazon Comprehend Medical, e.g. `POST /comprehendmedical/DetectEntitiesV2`.\n\nThe request body is forwarded as-is to the AWS JSON 1.1 API and signed with SigV4\nusing the proxy's AWS credentials.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/comprehend_medical)", + "operationId": "comprehend_medical_proxy_route_comprehendmedical__operation__post", + "parameters": [ + { + "in": "path", + "name": "operation", + "required": true, + "schema": { + "title": "Operation", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Comprehend Medical Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/cursor/chat/completions": { + "post": { + "description": "Cursor BYOK endpoint. Accepts both request shapes Cursor sends to its OpenAI-compatible\nbase URL and always answers in chat completions format.\n\nCursor agent mode sends Responses API format bodies (`input`, flat tool defs, `reasoning`,\ncustom tools) to the chat/completions path while expecting chat completions responses;\nthose are routed through the Responses API pipeline and converted back. Genuine chat\ncompletions bodies (`messages` present) are routed through the standard chat completions\npipeline, after normalizing each level of the `tools` array and `tool_choice` to the chat\ncompletions shapes OpenAI requires. Cursor mixes Responses API shapes into chat bodies\nper level, independently: a flat tool def (`{\"type\": \"custom\", \"name\": \"ApplyPatch\", ...}`)\ngets nested under `custom`, and a flat grammar format\n(`{\"type\": \"grammar\", \"definition\", \"syntax\"}`) gets wrapped as\n`{\"type\": \"grammar\", \"grammar\": {...}}` wherever it appears, including inside tool defs\nCursor already sent pre-nested.\n\n```bash\ncurl -X POST http://localhost:4000/cursor/chat/completions -H \"Content-Type: application/json\" -H \"Authorization: Bearer sk-1234\" -d '{\n \"model\": \"gpt-4o\",\n \"input\": [{\"role\": \"user\", \"content\": \"Hello\"}]\n}'\nResponds back in chat completions format.\n```", + "operationId": "cursor_chat_completions_cursor_chat_completions_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cursor Chat Completions", + "tags": [ + "llm_passthrough" + ] + } + }, + "/cursor/models": { + "get": { + "description": "OpenAI-compatible model listing for the Cursor BYOK base URL.\n\nClients pointed at `/cursor` as an OpenAI-compatible base URL resolve and\nverify models via `GET {base}/models` (the OpenAI SDK contract). Without this\nroute those requests fall through to the Cursor Cloud Agents passthrough, which\ndemands a Cursor API key and 401s, so key verification silently fails before any\nchat request is ever sent. Delegates to the standard `/v1/models` handler.", + "operationId": "cursor_model_list_cursor_models_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cursor Model List", + "tags": [ + "llm_passthrough" + ] + } + }, + "/cursor/v1/models": { + "get": { + "description": "OpenAI-compatible model listing for the Cursor BYOK base URL.\n\nClients pointed at `/cursor` as an OpenAI-compatible base URL resolve and\nverify models via `GET {base}/models` (the OpenAI SDK contract). Without this\nroute those requests fall through to the Cursor Cloud Agents passthrough, which\ndemands a Cursor API key and 401s, so key verification silently fails before any\nchat request is ever sent. Delegates to the standard `/v1/models` handler.", + "operationId": "cursor_model_list_cursor_v1_models_get", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cursor Model List", + "tags": [ + "llm_passthrough" + ] + } + }, + "/cursor/{endpoint}": { + "delete": { + "description": "Pass-through endpoint for the Cursor Cloud Agents API.\n\nSupports all Cursor Cloud Agents endpoints:\n- GET /v0/agents \u2014 List agents\n- POST /v0/agents \u2014 Launch an agent\n- GET /v0/agents/{id} \u2014 Agent status\n- GET /v0/agents/{id}/conversation \u2014 Agent conversation\n- POST /v0/agents/{id}/followup \u2014 Add follow-up\n- POST /v0/agents/{id}/stop \u2014 Stop an agent\n- DELETE /v0/agents/{id} \u2014 Delete an agent\n- GET /v0/me \u2014 API key info\n- GET /v0/models \u2014 List models\n- GET /v0/repositories \u2014 List GitHub repositories\n\nUses Basic Authentication (base64-encoded `API_KEY:`).\n\nCredential lookup order:\n1. passthrough_endpoint_router (config.yaml deployments with use_in_pass_through)\n2. litellm.credential_list (credentials added via UI)\n3. CURSOR_API_KEY environment variable", + "operationId": "cursor_proxy_route_cursor__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cursor Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Pass-through endpoint for the Cursor Cloud Agents API.\n\nSupports all Cursor Cloud Agents endpoints:\n- GET /v0/agents \u2014 List agents\n- POST /v0/agents \u2014 Launch an agent\n- GET /v0/agents/{id} \u2014 Agent status\n- GET /v0/agents/{id}/conversation \u2014 Agent conversation\n- POST /v0/agents/{id}/followup \u2014 Add follow-up\n- POST /v0/agents/{id}/stop \u2014 Stop an agent\n- DELETE /v0/agents/{id} \u2014 Delete an agent\n- GET /v0/me \u2014 API key info\n- GET /v0/models \u2014 List models\n- GET /v0/repositories \u2014 List GitHub repositories\n\nUses Basic Authentication (base64-encoded `API_KEY:`).\n\nCredential lookup order:\n1. passthrough_endpoint_router (config.yaml deployments with use_in_pass_through)\n2. litellm.credential_list (credentials added via UI)\n3. CURSOR_API_KEY environment variable", + "operationId": "cursor_proxy_route_cursor__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cursor Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Pass-through endpoint for the Cursor Cloud Agents API.\n\nSupports all Cursor Cloud Agents endpoints:\n- GET /v0/agents \u2014 List agents\n- POST /v0/agents \u2014 Launch an agent\n- GET /v0/agents/{id} \u2014 Agent status\n- GET /v0/agents/{id}/conversation \u2014 Agent conversation\n- POST /v0/agents/{id}/followup \u2014 Add follow-up\n- POST /v0/agents/{id}/stop \u2014 Stop an agent\n- DELETE /v0/agents/{id} \u2014 Delete an agent\n- GET /v0/me \u2014 API key info\n- GET /v0/models \u2014 List models\n- GET /v0/repositories \u2014 List GitHub repositories\n\nUses Basic Authentication (base64-encoded `API_KEY:`).\n\nCredential lookup order:\n1. passthrough_endpoint_router (config.yaml deployments with use_in_pass_through)\n2. litellm.credential_list (credentials added via UI)\n3. CURSOR_API_KEY environment variable", + "operationId": "cursor_proxy_route_cursor__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cursor Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Pass-through endpoint for the Cursor Cloud Agents API.\n\nSupports all Cursor Cloud Agents endpoints:\n- GET /v0/agents \u2014 List agents\n- POST /v0/agents \u2014 Launch an agent\n- GET /v0/agents/{id} \u2014 Agent status\n- GET /v0/agents/{id}/conversation \u2014 Agent conversation\n- POST /v0/agents/{id}/followup \u2014 Add follow-up\n- POST /v0/agents/{id}/stop \u2014 Stop an agent\n- DELETE /v0/agents/{id} \u2014 Delete an agent\n- GET /v0/me \u2014 API key info\n- GET /v0/models \u2014 List models\n- GET /v0/repositories \u2014 List GitHub repositories\n\nUses Basic Authentication (base64-encoded `API_KEY:`).\n\nCredential lookup order:\n1. passthrough_endpoint_router (config.yaml deployments with use_in_pass_through)\n2. litellm.credential_list (credentials added via UI)\n3. CURSOR_API_KEY environment variable", + "operationId": "cursor_proxy_route_cursor__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cursor Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Pass-through endpoint for the Cursor Cloud Agents API.\n\nSupports all Cursor Cloud Agents endpoints:\n- GET /v0/agents \u2014 List agents\n- POST /v0/agents \u2014 Launch an agent\n- GET /v0/agents/{id} \u2014 Agent status\n- GET /v0/agents/{id}/conversation \u2014 Agent conversation\n- POST /v0/agents/{id}/followup \u2014 Add follow-up\n- POST /v0/agents/{id}/stop \u2014 Stop an agent\n- DELETE /v0/agents/{id} \u2014 Delete an agent\n- GET /v0/me \u2014 API key info\n- GET /v0/models \u2014 List models\n- GET /v0/repositories \u2014 List GitHub repositories\n\nUses Basic Authentication (base64-encoded `API_KEY:`).\n\nCredential lookup order:\n1. passthrough_endpoint_router (config.yaml deployments with use_in_pass_through)\n2. litellm.credential_list (credentials added via UI)\n3. CURSOR_API_KEY environment variable", + "operationId": "cursor_proxy_route_cursor__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cursor Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/eu.assemblyai/{endpoint}": { + "delete": { + "operationId": "assemblyai_proxy_route_eu_assemblyai__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "operationId": "assemblyai_proxy_route_eu_assemblyai__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "operationId": "assemblyai_proxy_route_eu_assemblyai__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "operationId": "assemblyai_proxy_route_eu_assemblyai__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "operationId": "assemblyai_proxy_route_eu_assemblyai__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Assemblyai Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/gemini/{endpoint}": { + "delete": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio)", + "operationId": "gemini_proxy_route_gemini__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Gemini Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio)", + "operationId": "gemini_proxy_route_gemini__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Gemini Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio)", + "operationId": "gemini_proxy_route_gemini__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Gemini Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio)", + "operationId": "gemini_proxy_route_gemini__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Gemini Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/google_ai_studio)", + "operationId": "gemini_proxy_route_gemini__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Gemini Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/gigachat/{endpoint}": { + "delete": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/gigachat)", + "operationId": "gigachat_proxy_route_gigachat__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Gigachat Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/gigachat)", + "operationId": "gigachat_proxy_route_gigachat__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Gigachat Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/gigachat)", + "operationId": "gigachat_proxy_route_gigachat__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Gigachat Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/gigachat)", + "operationId": "gigachat_proxy_route_gigachat__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Gigachat Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/gigachat)", + "operationId": "gigachat_proxy_route_gigachat__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Gigachat Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/milvus/{endpoint}": { + "delete": { + "description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.", + "operationId": "milvus_proxy_route_milvus__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Milvus Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.", + "operationId": "milvus_proxy_route_milvus__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Milvus Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.", + "operationId": "milvus_proxy_route_milvus__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Milvus Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.", + "operationId": "milvus_proxy_route_milvus__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Milvus Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Enable using Milvus `/vectors` endpoint as a pass-through endpoint.", + "operationId": "milvus_proxy_route_milvus__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Milvus Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/mistral/{endpoint}": { + "delete": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/mistral)", + "operationId": "mistral_proxy_route_mistral__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Mistral Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/mistral)", + "operationId": "mistral_proxy_route_mistral__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Mistral Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/mistral)", + "operationId": "mistral_proxy_route_mistral__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Mistral Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/mistral)", + "operationId": "mistral_proxy_route_mistral__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Mistral Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/mistral)", + "operationId": "mistral_proxy_route_mistral__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Mistral Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/deployments/{model}/chat/completions": { + "post": { + "description": "Follows the exact same API spec as `OpenAI's Chat API https://platform.openai.com/docs/api-reference/chat`\n\n```bash\ncurl -X POST http://localhost:4000/v1/chat/completions \n-H \"Content-Type: application/json\" \n-H \"Authorization: Bearer sk-1234\" \n-d '{\n \"model\": \"gpt-4o\",\n \"messages\": [\n {\n \"role\": \"user\",\n \"content\": \"Hello!\"\n }\n ]\n}'\n```", + "operationId": "chat_completion_openai_deployments__model__chat_completions_post", + "parameters": [ + { + "in": "path", + "name": "model", + "required": true, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful response" + }, + "400": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "ContentPolicyViolationError" + }, + "401": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "AuthenticationError" + }, + "403": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "PermissionDeniedError" + }, + "404": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "NotFoundError" + }, + "408": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "Timeout" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "UnprocessableEntityError" + }, + "429": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "\n Unified rate-limit error.\n\n Every rate-limit condition surfaced by litellm \u2014 whether it originated from\n an upstream LLM provider, a vendor batch endpoint, or one of litellm's own\n proxy-side limiters (parallel-requests, dynamic-rate, batch-rate, budget,\n max-iterations, etc.) \u2014 is raised as an instance of this class.\n\n The :attr:`category` attribute lets callers distinguish the source. See\n :class:`RateLimitErrorCategory` for the available values.\n " + }, + "500": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "JSONSchemaValidationError" + }, + "503": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ErrorResponse" + } + } + }, + "description": "APIConnectionError" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Chat Completion", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/deployments/{model}/completions": { + "post": { + "description": "Follows the exact same API spec as `OpenAI's Completions API https://platform.openai.com/docs/api-reference/completions`\n\n```bash\ncurl -X POST http://localhost:4000/v1/completions \n-H \"Content-Type: application/json\" \n-H \"Authorization: Bearer sk-1234\" \n-d '{\n \"model\": \"gpt-3.5-turbo-instruct\",\n \"prompt\": \"Once upon a time\",\n \"max_tokens\": 50,\n \"temperature\": 0.7\n}'\n```", + "operationId": "completion_openai_deployments__model__completions_post", + "parameters": [ + { + "in": "path", + "name": "model", + "required": true, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Completion", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/deployments/{model}/embeddings": { + "post": { + "description": "Follows the exact same API spec as `OpenAI's Embeddings API https://platform.openai.com/docs/api-reference/embeddings`\n\n```bash\ncurl -X POST http://localhost:4000/v1/embeddings \n-H \"Content-Type: application/json\" \n-H \"Authorization: Bearer sk-1234\" \n-d '{\n \"model\": \"text-embedding-ada-002\",\n \"input\": \"The quick brown fox jumps over the lazy dog\"\n}'\n```", + "operationId": "embeddings_openai_deployments__model__embeddings_post", + "parameters": [ + { + "in": "path", + "name": "model", + "required": true, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Embeddings", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/deployments/{model}/images/edits": { + "post": { + "description": "Follows the OpenAI Images API spec: https://platform.openai.com/docs/api-reference/images/create\n\n```bash\ncurl -s -D >(grep -i x-request-id >&2) -o >(jq -r '.data[0].b64_json' | base64 --decode > gift-basket.png) -X POST \"http://localhost:4000/v1/images/edits\" -H \"Authorization: Bearer sk-1234\" -F \"model=gpt-image-1\" -F \"image[]=@soap.png\" -F 'prompt=Create a studio ghibli image of this'\n```", + "operationId": "image_edit_api_openai_deployments__model__images_edits_post", + "parameters": [ + { + "in": "path", + "name": "model", + "required": true, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + } + ], + "requestBody": { + "content": { + "multipart/form-data": { + "schema": { + "$ref": "#/components/schemas/Body_image_edit_api_openai_deployments__model__images_edits_post" + } + } + } + }, + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Image Edit Api", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/deployments/{model}/images/generations": { + "post": { + "operationId": "image_generation_openai_deployments__model__images_generations_post", + "parameters": [ + { + "in": "path", + "name": "model", + "required": true, + "schema": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Model" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Image Generation", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/realtime/calls": { + "post": { + "operationId": "proxy_realtime_calls_openai_v1_realtime_calls_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "summary": "Proxy Realtime Calls", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/realtime/client_secrets": { + "post": { + "operationId": "create_realtime_client_secret_openai_v1_realtime_client_secrets_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/RealtimeClientSecretResponse" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Create Realtime Client Secret", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/realtime/transcription_sessions": { + "post": { + "description": "Create an ephemeral Realtime transcription session\n(POST /v1/realtime/transcription_sessions) for the WebRTC/WebSocket flow.\n\nMirrors the client_secrets route but targets the transcription_sessions\nendpoint and encrypts the ephemeral key returned under `client_secret.value`.", + "operationId": "create_realtime_transcription_session_openai_v1_realtime_transcription_sessions_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/RealtimeTranscriptionSessionResponse" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Create Realtime Transcription Session", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/responses": { + "post": { + "description": "Follows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses\n\nSupports background mode with polling_via_cache for partial response retrieval.\nWhen background=true and polling_via_cache is enabled, returns a polling_id immediately\nand streams the response in the background, updating Redis cache.\n\n```bash\n# Normal request\ncurl -X POST http://localhost:4000/v1/responses -H \"Content-Type: application/json\" -H \"Authorization: Bearer sk-1234\" -d '{\n \"model\": \"gpt-4o\",\n \"input\": \"Tell me about AI\"\n}'\n\n# Background request with polling\ncurl -X POST http://localhost:4000/v1/responses -H \"Content-Type: application/json\" -H \"Authorization: Bearer sk-1234\" -d '{\n \"model\": \"gpt-4o\",\n \"input\": \"Tell me about AI\",\n \"background\": true\n}'\n```", + "operationId": "responses_api_openai_v1_responses_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Responses Api", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/responses/compact": { + "post": { + "description": "Compact a response by running a compaction pass over a conversation.\n\nReturns encrypted, opaque items that can be used to reduce context size.\n\nFollows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/compact\n\n```bash\ncurl -X POST http://localhost:4000/v1/responses/compact -H \"Content-Type: application/json\" -H \"Authorization: Bearer sk-1234\" -d '{\n \"model\": \"gpt-4o\",\n \"input\": [{\"role\": \"user\", \"content\": \"Hello\"}]\n}'\n```", + "operationId": "compact_response_openai_v1_responses_compact_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Compact Response", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/responses/input_tokens": { + "post": { + "description": "Count the input tokens of a Responses API request without calling the model.\n\nFollows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/input-tokens\n\n```bash\ncurl -X POST http://localhost:4000/v1/responses/input_tokens -H \"Content-Type: application/json\" -H \"Authorization: Bearer sk-1234\" -d '{\n \"model\": \"gpt-4o\",\n \"input\": \"Hello, how are you?\"\n}'\n```\n\nReturns: `{\"object\": \"response.input_tokens\", \"input_tokens\": }`", + "operationId": "responses_input_tokens_openai_v1_responses_input_tokens_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Responses Input Tokens", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/responses/{response_id}": { + "delete": { + "description": "Delete a response by ID.\n\nSupports both:\n- Polling IDs (litellm_poll_*): Deletes from Redis cache\n- Provider response IDs: Passes through to provider API\n\nFollows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/delete\n\n```bash\ncurl -X DELETE http://localhost:4000/v1/responses/resp_abc123 -H \"Authorization: Bearer sk-1234\"\n```", + "operationId": "delete_response_openai_v1_responses__response_id__delete", + "parameters": [ + { + "in": "path", + "name": "response_id", + "required": true, + "schema": { + "title": "Response Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Delete Response", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Get a response by ID.\n\nSupports both:\n- Polling IDs (litellm_poll_*): Returns cumulative cached content from background responses\n- Provider response IDs: Passes through to provider API\n\nFollows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/get\n\n```bash\n# Get polling response\ncurl -X GET http://localhost:4000/v1/responses/litellm_poll_abc123 -H \"Authorization: Bearer sk-1234\"\n\n# Get provider response\ncurl -X GET http://localhost:4000/v1/responses/resp_abc123 -H \"Authorization: Bearer sk-1234\"\n```", + "operationId": "get_response_openai_v1_responses__response_id__get", + "parameters": [ + { + "in": "path", + "name": "response_id", + "required": true, + "schema": { + "title": "Response Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Response", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/responses/{response_id}/cancel": { + "post": { + "description": "Cancel a response by ID.\n\nSupports both:\n- Polling IDs (litellm_poll_*): Cancels background response and updates status in Redis\n- Provider response IDs: Passes through to provider API\n\nFollows the OpenAI Responses API spec: https://platform.openai.com/docs/api-reference/responses/cancel\n\n```bash\n# Cancel polling response\ncurl -X POST http://localhost:4000/v1/responses/litellm_poll_abc123/cancel -H \"Authorization: Bearer sk-1234\"\n\n# Cancel provider response\ncurl -X POST http://localhost:4000/v1/responses/resp_abc123/cancel -H \"Authorization: Bearer sk-1234\"\n```", + "operationId": "cancel_response_openai_v1_responses__response_id__cancel_post", + "parameters": [ + { + "in": "path", + "name": "response_id", + "required": true, + "schema": { + "title": "Response Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Cancel Response", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/v1/responses/{response_id}/input_items": { + "get": { + "description": "List input items for a response.", + "operationId": "get_response_input_items_openai_v1_responses__response_id__input_items_get", + "parameters": [ + { + "in": "path", + "name": "response_id", + "required": true, + "schema": { + "title": "Response Id", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Get Response Input Items", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai/{endpoint}": { + "delete": { + "description": "Pass-through endpoint for OpenAI API calls.\n\nAvailable on both routes:\n- /openai/{endpoint:path} - Standard OpenAI passthrough route\n- /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)\n\nUse /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts\nwith LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).\n\nExamples:\n Standard route:\n - /openai/v1/chat/completions\n - /openai/v1/assistants\n - /openai/v1/threads\n\n Dedicated passthrough (for Responses API):\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_proxy_route_openai__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Pass-through endpoint for OpenAI API calls.\n\nAvailable on both routes:\n- /openai/{endpoint:path} - Standard OpenAI passthrough route\n- /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)\n\nUse /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts\nwith LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).\n\nExamples:\n Standard route:\n - /openai/v1/chat/completions\n - /openai/v1/assistants\n - /openai/v1/threads\n\n Dedicated passthrough (for Responses API):\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_proxy_route_openai__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Pass-through endpoint for OpenAI API calls.\n\nAvailable on both routes:\n- /openai/{endpoint:path} - Standard OpenAI passthrough route\n- /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)\n\nUse /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts\nwith LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).\n\nExamples:\n Standard route:\n - /openai/v1/chat/completions\n - /openai/v1/assistants\n - /openai/v1/threads\n\n Dedicated passthrough (for Responses API):\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_proxy_route_openai__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Pass-through endpoint for OpenAI API calls.\n\nAvailable on both routes:\n- /openai/{endpoint:path} - Standard OpenAI passthrough route\n- /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)\n\nUse /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts\nwith LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).\n\nExamples:\n Standard route:\n - /openai/v1/chat/completions\n - /openai/v1/assistants\n - /openai/v1/threads\n\n Dedicated passthrough (for Responses API):\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_proxy_route_openai__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Pass-through endpoint for OpenAI API calls.\n\nAvailable on both routes:\n- /openai/{endpoint:path} - Standard OpenAI passthrough route\n- /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API)\n\nUse /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts\nwith LiteLLM's native implementations (e.g., for the Responses API at /v1/responses).\n\nExamples:\n Standard route:\n - /openai/v1/chat/completions\n - /openai/v1/assistants\n - /openai/v1/threads\n\n Dedicated passthrough (for Responses API):\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_proxy_route_openai__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/openai_passthrough/{endpoint}": { + "delete": { + "description": "Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native\nimplementations (e.g. the Responses API at /v1/responses).\n\nExamples:\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_passthrough_route_openai_passthrough__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Passthrough Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native\nimplementations (e.g. the Responses API at /v1/responses).\n\nExamples:\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_passthrough_route_openai_passthrough__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Passthrough Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native\nimplementations (e.g. the Responses API at /v1/responses).\n\nExamples:\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_passthrough_route_openai_passthrough__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Passthrough Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native\nimplementations (e.g. the Responses API at /v1/responses).\n\nExamples:\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_passthrough_route_openai_passthrough__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Passthrough Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native\nimplementations (e.g. the Responses API at /v1/responses).\n\nExamples:\n - /openai_passthrough/v1/responses\n - /openai_passthrough/v1/responses/{response_id}\n - /openai_passthrough/v1/responses/{response_id}/input_items\n\n[Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough)", + "operationId": "openai_passthrough_route_openai_passthrough__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Openai Passthrough Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/vertex_ai/discovery/{endpoint}": { + "delete": { + "description": "Call any vertex discovery endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/vertex_ai/discovery/{endpoint:path}`\n\nTarget url: `https://discoveryengine.googleapis.com`", + "operationId": "vertex_discovery_proxy_route_vertex_ai_discovery__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Vertex Discovery Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Call any vertex discovery endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/vertex_ai/discovery/{endpoint:path}`\n\nTarget url: `https://discoveryengine.googleapis.com`", + "operationId": "vertex_discovery_proxy_route_vertex_ai_discovery__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Vertex Discovery Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Call any vertex discovery endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/vertex_ai/discovery/{endpoint:path}`\n\nTarget url: `https://discoveryengine.googleapis.com`", + "operationId": "vertex_discovery_proxy_route_vertex_ai_discovery__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Vertex Discovery Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Call any vertex discovery endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/vertex_ai/discovery/{endpoint:path}`\n\nTarget url: `https://discoveryengine.googleapis.com`", + "operationId": "vertex_discovery_proxy_route_vertex_ai_discovery__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Vertex Discovery Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Call any vertex discovery endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/vertex_ai/discovery/{endpoint:path}`\n\nTarget url: `https://discoveryengine.googleapis.com`", + "operationId": "vertex_discovery_proxy_route_vertex_ai_discovery__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "summary": "Vertex Discovery Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/vertex_ai/{endpoint}": { + "delete": { + "description": "Call LiteLLM proxy via Vertex AI SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)", + "operationId": "vertex_proxy_route_vertex_ai__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vertex Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Call LiteLLM proxy via Vertex AI SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)", + "operationId": "vertex_proxy_route_vertex_ai__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vertex Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Call LiteLLM proxy via Vertex AI SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)", + "operationId": "vertex_proxy_route_vertex_ai__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vertex Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Call LiteLLM proxy via Vertex AI SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)", + "operationId": "vertex_proxy_route_vertex_ai__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vertex Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Call LiteLLM proxy via Vertex AI SDK.\n\n[Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai)", + "operationId": "vertex_proxy_route_vertex_ai__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vertex Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/vllm/{endpoint}": { + "delete": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/vllm)", + "operationId": "vllm_proxy_route_vllm__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vllm Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/vllm)", + "operationId": "vllm_proxy_route_vllm__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vllm Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/vllm)", + "operationId": "vllm_proxy_route_vllm__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vllm Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/vllm)", + "operationId": "vllm_proxy_route_vllm__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vllm Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "[Docs](https://docs.litellm.ai/docs/pass_through/vllm)", + "operationId": "vllm_proxy_route_vllm__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Vllm Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + }, + "/watsonx/{endpoint}": { + "delete": { + "description": "Watsonx pass-through endpoint.\nAllows using Watsonx APIs with automatic IAM token management and version parameter injection.\n\nExample:\n POST /watsonx/ml/v1/text/tokenization\n POST /watsonx/ml/v1/text/generation", + "operationId": "watsonx_proxy_route_watsonx__endpoint__delete", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Watsonx Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "get": { + "description": "Watsonx pass-through endpoint.\nAllows using Watsonx APIs with automatic IAM token management and version parameter injection.\n\nExample:\n POST /watsonx/ml/v1/text/tokenization\n POST /watsonx/ml/v1/text/generation", + "operationId": "watsonx_proxy_route_watsonx__endpoint__get", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Watsonx Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "patch": { + "description": "Watsonx pass-through endpoint.\nAllows using Watsonx APIs with automatic IAM token management and version parameter injection.\n\nExample:\n POST /watsonx/ml/v1/text/tokenization\n POST /watsonx/ml/v1/text/generation", + "operationId": "watsonx_proxy_route_watsonx__endpoint__patch", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Watsonx Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "post": { + "description": "Watsonx pass-through endpoint.\nAllows using Watsonx APIs with automatic IAM token management and version parameter injection.\n\nExample:\n POST /watsonx/ml/v1/text/tokenization\n POST /watsonx/ml/v1/text/generation", + "operationId": "watsonx_proxy_route_watsonx__endpoint__post", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Watsonx Proxy Route", + "tags": [ + "llm_passthrough" + ] + }, + "put": { + "description": "Watsonx pass-through endpoint.\nAllows using Watsonx APIs with automatic IAM token management and version parameter injection.\n\nExample:\n POST /watsonx/ml/v1/text/tokenization\n POST /watsonx/ml/v1/text/generation", + "operationId": "watsonx_proxy_route_watsonx__endpoint__put", + "parameters": [ + { + "in": "path", + "name": "endpoint", + "required": true, + "schema": { + "title": "Endpoint", + "type": "string" + } + } + ], + "responses": { + "200": { + "content": { + "application/json": { + "schema": {} + } + }, + "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Watsonx Proxy Route", + "tags": [ + "llm_passthrough" + ] + } + } + } + }, "mcp_app": { "components": { "schemas": { @@ -31965,7 +36935,7 @@ "paths": { "/openai/v1/realtime/calls": { "post": { - "operationId": "proxy_realtime_calls_openai_v1_realtime_calls_post", + "operationId": "proxy_realtime_calls_openai_v1_realtime_calls_post_2", "responses": { "200": { "content": { @@ -31984,7 +36954,7 @@ }, "/openai/v1/realtime/client_secrets": { "post": { - "operationId": "create_realtime_client_secret_openai_v1_realtime_client_secrets_post", + "operationId": "create_realtime_client_secret_openai_v1_realtime_client_secrets_post_2", "responses": { "200": { "content": { @@ -32011,7 +36981,7 @@ "/openai/v1/realtime/transcription_sessions": { "post": { "description": "Create an ephemeral Realtime transcription session\n(POST /v1/realtime/transcription_sessions) for the WebRTC/WebSocket flow.\n\nMirrors the client_secrets route but targets the transcription_sessions\nendpoint and encrypts the ephemeral key returned under `client_secret.value`.", - "operationId": "create_realtime_transcription_session_openai_v1_realtime_transcription_sessions_post", + "operationId": "create_realtime_transcription_session_openai_v1_realtime_transcription_sessions_post_2", "responses": { "200": { "content": { diff --git a/litellm/proxy/_types.py b/litellm/proxy/_types.py index 2cbff128635..7b0e92d81ae 100644 --- a/litellm/proxy/_types.py +++ b/litellm/proxy/_types.py @@ -880,6 +880,8 @@ class LiteLLMRoutes(enum.Enum): # proxy admin, or team admin naming their own team via team_id "/auto_router/test_routing", "/auto_router/validate_complexity_router_config", + # Per-session auto-router read - the endpoint scopes the row to the caller's own key hash + "/auto_router/session", # Agent registry - reads are role-scoped and writes are proxy-admin-gated # inside agent_endpoints/endpoints.py *agent_management_routes, @@ -2602,6 +2604,15 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): global_max_parallel_requests: int | None = Field( None, description="global max parallel requests to allow for a proxy instance." ) + user_api_key_cache_max_size: int | None = Field( + None, + gt=0, + description=( + "max number of entries (virtual keys, teams, users, end users, memberships, ...) each worker keeps in " + "its in-memory auth cache. Defaults to 200. Raise this if you have more active keys than that or auth " + "lookups keep hitting the DB" + ), + ) max_request_size_mb: int | None = Field( None, description="max request size in MB, if a request is larger than this size it will be rejected", @@ -2848,6 +2859,25 @@ class ConfigGeneralSettings(LiteLLMPydanticObjectBase): "UI username/password login. Default is False." ), ) + disable_responses_id_security: bool | None = Field( + None, + description=( + "If True, disables ownership enforcement on Responses API ids. " + "Keys may then retrieve, cancel, delete, and chain from any response id, " + "including ids belonging to another user or team and ids this proxy never issued. " + "WARNING: this removes tenant isolation on /v1/responses" + ), + ) + allow_unmanaged_response_ids: bool | None = Field( + None, + description=( + "If True, lets keys address Responses API ids that this proxy did not issue " + "(raw provider ids, or ids issued before response-id encryption was configured). " + "Such an id carries no owner, so no ownership check can run on it; ids this proxy " + "did issue keep full ownership enforcement. Off by default, in which case an " + "unrecognized response id is rejected with 403" + ), + ) disable_env_credential_login: bool | None = Field( None, description=( diff --git a/litellm/proxy/client/cli/README.md b/litellm/proxy/client/cli/README.md index 4045857a237..2071576a943 100644 --- a/litellm/proxy/client/cli/README.md +++ b/litellm/proxy/client/cli/README.md @@ -508,9 +508,9 @@ The credential is short-lived by design (default 24h, configurable via `LITELLM_ ### Route Every Claude Code Session Through the Proxy -`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` to `1` when those keys are missing, and `apiKeyHelper` to a `lite auth print-token` invocation, drops any stray static `ANTHROPIC_API_KEY` or `ANTHROPIC_AUTH_TOKEN` so the helper-issued token wins, and leaves every other setting in the file untouched. It backs up the original file before patching it. +`lite claude` wraps a single invocation, but `lite up` goes further: it patches `~/.claude/settings.json`, Claude Code's own config file, so that every Claude Code session started afterward -- from any terminal, launched normally with just `claude`, no wrapper needed -- routes through your LiteLLM proxy. It sets `env.ANTHROPIC_BASE_URL` to the proxy URL, `env.ENABLE_TOOL_SEARCH` to `true` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` to `1` when those keys are missing, writes the key it resolved (your fresh `lite login`, or an explicit `--api-key`) into `env.ANTHROPIC_AUTH_TOKEN` as a static token, drops any stray `ANTHROPIC_API_KEY` or `apiKeyHelper` so nothing fights that token, and leaves every other setting in the file untouched. It backs up the original file before patching it. Nothing here writes an `apiKeyHelper`: Claude Code would spawn `lite` (and its keychain check) on every credential refresh, so the key is copied in instead and `lite up` restores the file when it stops. -Two things need to already be true: you've run `lite login` (or `lite login --pkce`, whose key the helper renews on its own), since the apiKeyHelper depends on that stored token, and the proxy is already reachable, since `lite up` does not start one for you. +Two things need to already be true: you've run `lite login` (or passed a key), and the proxy is already reachable, since `lite up` does not start one for you. ```bash lite login @@ -526,21 +526,19 @@ Cursor is not supported: it has no equivalent file-based config to hot-patch thi #### Making It Permanent at Login -`lite up` holds the patch only for as long as it runs. To wire Claude Code up once and leave it that way, pass `--config-claude` to `lite login`: +`lite up` holds its patch only for as long as it runs. To wire Claude Code up at login and leave it that way, pass `--config-claude` to `lite login`: ```bash lite --base-url https://your-proxy.example.com login --config-claude ``` -It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY`, and `apiKeyHelper`, but persistently: no foreground process to keep alive, and `lite unconfigure claude` restores what it changed (see below). Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag. +It writes the same settings `lite up` does, `env.ANTHROPIC_BASE_URL`, `env.ENABLE_TOOL_SEARCH`, `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY`, and the key this login minted as `env.ANTHROPIC_AUTH_TOKEN`, but persistently: no foreground process to keep alive, and `lite unconfigure claude` restores what it changed (see below). Every other key in `~/.claude/settings.json` is preserved, the file is created if it does not exist, and it is written atomically with owner-only permissions. Plain `lite login` is unchanged; nothing happens to your Claude Code config unless you pass the flag -Because the credential is reached through `apiKeyHelper` rather than copied into the file, a later `lite login` refreshes it with no further action: Claude Code re-runs the helper on every request and picks up whatever token the most recent login stored. Nothing secret is written to `settings.json`. +The key in the file is the login's own, so it expires with it (24h by default): run `lite login --config-claude` again after that, which rewrites the key in place. Earlier versions wrote an `apiKeyHelper` that ran `lite auth print-token` instead, so a later login refreshed Claude Code by itself; that meant Claude Code spawning a full `lite` start, keychain check included, on every credential refresh, so the helper is no longer written and a stale one is stripped by the next `--config-claude` or `configure claude`. Like `lite up`, the flag refuses to run while a `lite up` session holds a backup, and tells you to run `lite down` first -Run it again to point Claude Code at a different proxy; the base URL and the helper are both rewritten. `lite up` and `--config-claude` manage the same file, so the flag refuses to run while a `lite up` session holds a backup, and tells you to run `lite down` first, rather than writing settings that `lite up` would silently revert when it stops. +#### Configuring Claude Code Once, With a Virtual Key -#### Configuring Claude Code Once, With a Virtual Key or Your Login - -`lite configure claude` wires Claude Code up persistently and `lite unconfigure claude` puts things back. It is what `lite login --config-claude` does, plus a pinned model and an undo, and it also takes a long-lived virtual key when that is what you have: +`lite configure claude` wires Claude Code up persistently with a long-lived virtual key, a pinned model and an undo, and `lite unconfigure claude` puts things back: ```bash curl -fsSL https://raw.githubusercontent.com/BerriAI/litellm/main/scripts/install.sh | sh @@ -548,11 +546,25 @@ lite --base-url https://your-proxy.example.com configure claude --api-key sk-... claude ``` -With `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`) the key is written into `env.ANTHROPIC_AUTH_TOKEN`. Without one, your `lite login` credential is used the way `--config-claude` uses it, through `apiKeyHelper`, so a later `lite login` (or a `--pkce` renewal) picks up on its own and nothing secret lands in the file; a missing or stale login is refreshed first. Either way the command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (under `claude-router-` for a group whose id contains neither `claude` nor `anthropic`, since Claude Code lists only those) and you pick between them as usual. Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key, which has to be on `/v1/models` for the key. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute up` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window (the proxy appends `[1m]` for a group whose configured or known input window reaches 1M) and sends no thinking parameters for it, so name the group like a Claude model id to change that. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control +The key comes from `--api-key` (or `lite --api-key` / `LITELLM_PROXY_API_KEY`) and is written into `env.ANTHROPIC_AUTH_TOKEN`; without one the command refuses, since a `lite login` credential expires within a day and keeping it fresh would mean Claude Code running `lite` through `apiKeyHelper` on every credential refresh. The command checks the key against `GET /v1/models`, then patches `~/.claude/settings.json`: `env.ANTHROPIC_BASE_URL`, the credential, and `env.ENABLE_TOOL_SEARCH` and `env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY` when those are missing, so Claude Code's `/model` picker lists the proxy's models (under `claude-router-` for a group whose id contains neither `claude` nor `anthropic`, since Claude Code lists only those) and you pick between them as usual. Claude Code keeps its own default model until you switch, so that id has to exist on the proxy for the first message to go through; `--model` (or the interactive prompt below) sets the model Claude Code starts on instead, as the top-level `model` key and as `env.ANTHROPIC_MODEL`, both of which have to be on `/v1/models` for the key. The second one matters for `claude -c` and `claude --resume`: a resumed session otherwise re-sends the model its transcript recorded, which behind an auto-router with `return_raw_model_name: true` is the tier model that answered, and a key scoped to the router alias gets a 403 for it; `ANTHROPIC_MODEL` outranks the transcript on resume. Nothing forces Claude Code's sub-agent or background tiers onto a proxy model, so those built-in ids need to exist on the proxy too; `lite autoroute up` is the mode that pins every tier to one group. Claude Code treats a name it does not know as an unknown model: it prints a one-line `unrecognized_model` note, assumes a 200k context window (the proxy appends `[1m]` for a group whose configured or known input window reaches 1M) and sends no thinking parameters for it, so name the group like a Claude model id to change that. The other credential slots (`env.ANTHROPIC_API_KEY`, a stale `env.ANTHROPIC_AUTH_TOKEN` or `apiKeyHelper`) are removed so they cannot fight the one written. Every other setting is preserved and the file is written atomically with owner-only permissions; if `settings.json` is a symlink into a dotfiles repository, the key is written through to that target and the command says so, so keep it out of version control Plain `lite configure`, with no agent named, asks the same things interactively: which agents to wire (Claude Code today) and which of the proxy's models to start on, picked from `/v1/models` with a type-to-filter prompt -What the command changed is recorded in `~/.litellm/claude_configure_state.json` (previous values plus fingerprints of what was written, never a second copy of the key). `lite unconfigure claude` restores each of those keys only if it still holds what `configure` wrote, so anything you changed since is left alone and named in the output; a `settings.json` or `env` object that only existed because of `configure` is removed again. Ownership moves only by a write: running `configure` again (a re-login is one) refreshes the record only for the keys its merge changed, keeps the original snapshot of a key that still holds what it wrote, and snapshots afresh a key you changed in between, so `unconfigure` brings back whatever the repeat displaced and never adopts your edit as its own. A credential (`env.ANTHROPIC_API_KEY`, `env.ANTHROPIC_AUTH_TOKEN`, `apiKeyHelper`) is put back only when the restored file points at the `ANTHROPIC_BASE_URL` it was captured next to; otherwise it stays removed, the output says which server it belonged to, and the receipt is kept so pointing the URL back and running `unconfigure` again finishes the job. It also undoes `lite login --config-claude`, which writes through the same path. Like `--config-claude`, both refuse to run while a `lite up` or `lite autoroute up` session holds a backup, and that check comes before any login prompt or request +What the command changed is recorded in `~/.litellm/claude_configure_state.json` (previous values plus fingerprints of what was written, never a second copy of the key). `lite unconfigure claude` restores each of those keys only if it still holds what `configure` wrote, so anything you changed since is left alone and named in the output; a `settings.json` or `env` object that only existed because of `configure` is removed again. Ownership moves only by a write: running `configure` again (a re-login is one) refreshes the record only for the keys its merge changed, keeps the original snapshot of a key that still holds what it wrote, and snapshots afresh a key you changed in between, so `unconfigure` brings back whatever the repeat displaced and never adopts your edit as its own. A credential (`env.ANTHROPIC_API_KEY`, `env.ANTHROPIC_AUTH_TOKEN`, `apiKeyHelper`) is put back only when the restored file points at the `ANTHROPIC_BASE_URL` it was captured next to; otherwise it stays removed, the output says which server it belonged to, and the receipt is kept so pointing the URL back and running `unconfigure` again finishes the job. It also undoes `lite login --config-claude`, which writes through the same path. Both refuse to run while a `lite up` or `lite autoroute up` session holds a backup, and that check comes before any request + +#### Routed model and savings in the status line + +`lite configure claude`, `lite login --config-claude`, `lite up` and `lite autoroute up` also install a status line (`~/.litellm/statusline.py`, registered as `statusLine` in `~/.claude/settings.json` unless you already run one) that shows which model the auto-router actually served the last turn and, once the proxy has recorded the session, what the session cost against the router's savings baseline: + +``` +claude-auto · Routed to: claude-haiku-4-5 -63% vs Claude Opus 5 +LiteLLM ████████░░░░░░░░░░░░░░░░ $0.14 +Claude Opus 5 ████████████████████████ $0.38 +``` + +The routed model comes from Claude Code's own transcript, so it only names the tier model when the auto-router deployment sets `return_raw_model_name: true` (the `lite autoroute` wizard does); otherwise it shows the alias you requested. The cost lines come from `GET /auto_router/session?session_id=...`, which any virtual key may call for its own sessions, and are cached for five seconds under a per-user `$TMPDIR/litellm-statusline-` directory. The baseline is the priciest model in the router's hardest tier, the same counterfactual the auto-router's savings reports use. `lite unconfigure claude` removes the `statusLine` entry only while it still points at that script. + +`lite codex` registers the same script as a Codex `Stop` hook for the launch, so after each turn Codex prints the same block as a system message. Codex asks once to trust the hook; the answer is remembered for later launches. ### QA Complexity-Based Auto-Routing Against Your Real Proxy diff --git a/litellm/proxy/client/cli/commands/agents.py b/litellm/proxy/client/cli/commands/agents.py index bf2a784f590..ea1eed65505 100644 --- a/litellm/proxy/client/cli/commands/agents.py +++ b/litellm/proxy/client/cli/commands/agents.py @@ -1,4 +1,6 @@ +import json import os +import re import shutil import subprocess import sys @@ -13,7 +15,7 @@ import requests from pydantic import BaseModel, TypeAdapter, ValidationError from .auth import CliContextObj, context_secret_vault, get_stored_api_key, login -from .claude_settings import claude_settings_path, lite_api_key_helper_configured +from .claude_settings import ClaudeSettingsError, install_statusline_script from .cmd_quoting import quote_for_cmd from .pi import ( LITELLM_PROXY_API_KEY_ENV, @@ -86,8 +88,6 @@ def build_agent_env( base_url: str, api_key: str, profiles: frozenset[str], - *, - export_anthropic_token: bool = True, ) -> dict[str, str]: """Return a copy of base_env wired to route the agent through the proxy. @@ -102,19 +102,12 @@ def build_agent_env( proxy's /v1/models; likewise left alone when already set. pi ignores both base URL variables and instead resolves $LITELLM_PROXY_API_KEY from its synced models.json provider entry. - - With export_anthropic_token=False the bearer is left out (and any inherited - one dropped) so Claude Code asks its configured apiKeyHelper instead; Claude - Code prefers ANTHROPIC_AUTH_TOKEN over the helper and warns when both are set. """ env: Final = dict(base_env) root: Final = base_url.rstrip("/") if PROFILE_ANTHROPIC in profiles: env[ANTHROPIC_BASE_URL_ENV] = root - if export_anthropic_token: - env[ANTHROPIC_AUTH_TOKEN_ENV] = api_key - else: - env.pop(ANTHROPIC_AUTH_TOKEN_ENV, None) + env[ANTHROPIC_AUTH_TOKEN_ENV] = api_key env.pop(ANTHROPIC_API_KEY_ENV, None) if ENABLE_TOOL_SEARCH_ENV not in env: env[ENABLE_TOOL_SEARCH_ENV] = ENABLE_TOOL_SEARCH_VALUE @@ -189,10 +182,52 @@ def prepare_pi( return ("--model", f"{PI_PROVIDER_NAME}/{ids[0]}") +def _warn(message: str) -> None: + click.echo(message, err=True) + + +_CODEX_STOP_HOOKS_DECLARED: Final = re.compile( + r"^\s*(\[\[\s*\"?hooks\"?\s*\.\s*\"?Stop\"?\s*\]\]|\"?hooks\"?(?:\s*\.\s*\"?Stop\"?)?\s*=|\[\s*\"?hooks\"?\s*\])", + re.MULTILINE, +) + + +def codex_config_path(base_env: Mapping[str, str]) -> Path: + return Path(base_env.get("CODEX_HOME") or Path.home() / ".codex") / "config.toml" + + +def codex_declares_stop_hooks(config_path: Path) -> bool: + """A config that cannot be read or decoded declares nothing we can see; Codex reports its own + TOML failure at launch, so the pre-check must not be the thing that stops `lite codex`.""" + try: + return _CODEX_STOP_HOOKS_DECLARED.search(config_path.read_text(encoding="utf-8")) is not None + except (OSError, UnicodeDecodeError): + return False + + +def prepare_codex( + base_url: str, + api_key: str, + base_env: Mapping[str, str], + *, + install: Callable[[], str] = install_statusline_script, + warn: Callable[[str], None] = _warn, +) -> tuple[str, ...]: + """A `-c hooks.Stop=` session flag replaces the user's whole Stop list, so their own hooks win over ours.""" + if codex_declares_stop_hooks(codex_config_path(base_env)): + warn("litellm: your Codex config already declares hooks; not adding the routed-model Stop hook") + return () + try: + command: Final = install() + except ClaudeSettingsError as e: + raise AgentRunError(str(e)) from e + return ("-c", f'hooks.Stop=[{{hooks=[{{type="command",command={json.dumps(command)}}}]}}]') + + _Preparer: TypeAlias = Callable[[str, str, Mapping[str, str]], Sequence[str]] _PREPARERS: Final[Mapping[str, _Preparer]] = MappingProxyType( - {"pi": prepare_pi} # mutable-ok: MappingProxyType freezes the provider registry + {"pi": prepare_pi, "codex": prepare_codex} # mutable-ok: MappingProxyType freezes the provider registry ) @@ -454,10 +489,6 @@ def _restore_controlling_terminal() -> None: os.close(fd) -def _warn(message: str) -> None: - click.echo(message, err=True) - - def run_agent( base_url: str, api_key: str, @@ -474,7 +505,6 @@ def run_agent( launcher: Callable[[str, Sequence[str], Mapping[str, str]], None] = _hand_off, reattach_terminal: Callable[[], None] | None = None, preparers: Mapping[str, _Preparer] = MappingProxyType(_PREPARERS), - export_anthropic_token: bool = True, ) -> None: """Validate, wire the environment, and hand off to the agent. @@ -506,9 +536,7 @@ def run_agent( env: Final = MappingProxyType( { - **build_agent_env( - env_before_sync, base_url, api_key, profiles, export_anthropic_token=export_anthropic_token - ), + **build_agent_env(env_before_sync, base_url, api_key, profiles), **(_NO_EXTRA_ENV if isinstance(synced, ModelSyncSkipped) else synced), } ) @@ -546,26 +574,14 @@ def resolve_api_key(ctx: click.Context) -> str: _SKIP_VERIFY_HELP: Final = "Skip the pre-launch key check against the proxy." -def _helper_supplies_token( - ctx_obj: CliContextObj, base_url: str, profiles: frozenset[str], settings_path: Path -) -> bool: - if PROFILE_ANTHROPIC not in profiles or not ctx_obj.get("api_key_from_token_file"): - return False - return lite_api_key_helper_configured(base_url, settings_path) - - def _launch(ctx: click.Context, binary: str, args: Sequence[str], *, skip_verify: bool) -> None: ctx_obj: Final[CliContextObj] = ctx.obj base_url: Final = ctx_obj["base_url"] started_interactive: Final = _is_interactive() api_key: Final = resolve_api_key(ctx) - display_name, profiles = agent_profile(binary) - settings_path: Final = claude_settings_path(os.environ) - helper_supplies_token: Final = _helper_supplies_token(ctx_obj, base_url, profiles, settings_path) + display_name, _profiles = agent_profile(binary) click.echo(f"litellm: routing {display_name} through proxy at {base_url.rstrip('/')}") - if helper_supplies_token: - click.echo(f"litellm: {display_name} reads its key from the apiKeyHelper in {settings_path}") try: run_agent( @@ -574,7 +590,6 @@ def _launch(ctx: click.Context, binary: str, args: Sequence[str], *, skip_verify [binary, *args], skip_verify=skip_verify, reattach_terminal=(_restore_controlling_terminal if started_interactive else None), - export_anthropic_token=not helper_supplies_token, ) except AgentRunError as e: raise click.ClickException(str(e)) diff --git a/litellm/proxy/client/cli/commands/auth.py b/litellm/proxy/client/cli/commands/auth.py index 4f704afe9d6..98af32fa7aa 100644 --- a/litellm/proxy/client/cli/commands/auth.py +++ b/litellm/proxy/client/cli/commands/auth.py @@ -42,14 +42,13 @@ from litellm.litellm_core_utils.cli_token_utils import ( from .claude_settings import ( STARTING_MODEL_ROLE, - ApiKeyHelper, ClaudeSettingsError, KeepModel, + StaticToken, claude_settings_path, configure_claude_settings, configure_state_path, refuse_while_owned, - resolve_api_key_helper, settings_file_owners, ) from .pkce_login import ( @@ -784,13 +783,16 @@ def _render_and_prompt_for_team_selection(teams: list[CliTeam]) -> str | None: return None -def _configure_claude_code(base_url: str) -> None: - """Point Claude Code at base_url by patching the settings.json it reads, undoable with `lite unconfigure claude`.""" +def _configure_claude_code(base_url: str, api_key: str) -> None: + """Write the key this login just minted into Claude Code's settings.json as a static token, undoable with + `lite unconfigure claude`. The key expires with the login, so the flag is the re-wire step of each login + rather than a one-time setup: no apiKeyHelper is written, since Claude Code would spawn `lite` (and its + keychain probe) on every credential refresh to keep one fresh.""" settings_path: Final = claude_settings_path(os.environ) try: configure_claude_settings( base_url, - ApiKeyHelper(resolve_api_key_helper(base_url)), + StaticToken(api_key), KeepModel(), settings_path, configure_state_path(settings_path), @@ -800,22 +802,28 @@ def _configure_claude_code(base_url: str) -> None: raise click.ClickException(f"Logged in, but could not configure Claude Code: {e}") click.echo(f"\nConfigured Claude Code: {settings_path} now routes through {base_url.rstrip('/')}.") click.echo( + "This login's key is stored in the file, so run `lite login --config-claude` again after it expires. " "Your other Claude Code settings were left untouched. Restart Claude Code to pick this up. " f"Undo with `lite unconfigure claude`; `lite configure claude --model` sets {STARTING_MODEL_ROLE}." ) def _finish_login(base_url: str, api_key: str, config_claude: bool, stored: SecretSave) -> None: + """Claude Code is configured from the key in hand, so it does not wait on the CLI's own store: a login whose + token file or keychain refused it still has a usable key, and `--config-claude` asked for exactly that + key to be written into settings.json.""" from litellm.proxy.client.cli.interface import show_commands click.echo("\nLogin successful!") click.echo(f"JWT Token: {api_key[:20]}...") click.echo(storage_notice(stored)) + if config_claude: + _configure_claude_code(base_url, api_key) if isinstance(stored, (CredentialNotSaved, CredentialNotRecorded)): + if config_claude: + click.echo("Claude Code was configured with this key even though the CLI itself could not keep it.") return click.echo("You can now use the CLI without specifying --api-key") - if config_claude: - _configure_claude_code(base_url) click.echo("\n" + "=" * 60) show_commands() @@ -850,8 +858,8 @@ def _pkce_login(base_url: str, config_claude: bool, vault: SecretVault) -> None: is_flag=True, default=False, help=( - "After logging in, update ~/.claude/settings.json so Claude Code routes through this proxy. " - "Unrelated settings are preserved." + "After logging in, write this login's key into ~/.claude/settings.json so Claude Code routes through " + "this proxy; run it again after the key expires. Unrelated settings are preserved." ), ) @click.option( diff --git a/litellm/proxy/client/cli/commands/autoroute/commands.py b/litellm/proxy/client/cli/commands/autoroute/commands.py index 86701f186fd..5d91fc81350 100644 --- a/litellm/proxy/client/cli/commands/autoroute/commands.py +++ b/litellm/proxy/client/cli/commands/autoroute/commands.py @@ -1,5 +1,4 @@ import atexit -import json import secrets import signal import threading @@ -15,8 +14,10 @@ from ..claude_settings import ( CLAUDE_SETTINGS_PATH, ClaudeSettingsError, StaticToken, + install_statusline_script, load_json_or_empty, merge_claude_settings, + write_claude_settings, ) from ..up import BackupRecord as ClaudeBackupRecord from ..up import restore_claude_settings, write_backup @@ -151,6 +152,7 @@ def up(port: int) -> None: raise click.ClickException(str(e)) try: + status_line: Final = install_statusline_script() original_existed: Final = CLAUDE_SETTINGS_PATH.exists() original_settings: Final = load_json_or_empty(CLAUDE_SETTINGS_PATH) write_backup( @@ -158,11 +160,15 @@ def up(port: int) -> None: AUTOROUTE_BACKUP_PATH, ) merged: Final = merge_claude_settings( - original_settings, base_url, StaticToken(master_key), AUTOROUTER_MODEL_NAME, AUTOROUTER_MODEL_NAME + original_settings, + base_url, + StaticToken(master_key), + AUTOROUTER_MODEL_NAME, + AUTOROUTER_MODEL_NAME, + status_line=status_line, ) CLAUDE_SETTINGS_PATH.parent.mkdir(parents=True, exist_ok=True) - with secure_create(CLAUDE_SETTINGS_PATH) as f: - json.dump(merged, f, indent=2) + write_claude_settings(CLAUDE_SETTINGS_PATH, merged) except ClaudeSettingsError as e: terminate(process.pid) clear_pid_record() diff --git a/litellm/proxy/client/cli/commands/autoroute/config.py b/litellm/proxy/client/cli/commands/autoroute/config.py index 9dfc4ad079b..1f3ad34e3d9 100644 --- a/litellm/proxy/client/cli/commands/autoroute/config.py +++ b/litellm/proxy/client/cli/commands/autoroute/config.py @@ -162,6 +162,7 @@ def build_generated_model_list(config: AutorouteConfig) -> list[JsonValue]: complexity_router_config: Final[dict[str, JsonValue]] = { "tiers": {tier: list(models) for tier, models in config.tiers.items()}, "default_model": config.default_model, + "return_raw_model_name": True, } if isinstance(config.classifier, LLMClassifier): complexity_router_config["classifier_type"] = "llm" diff --git a/litellm/proxy/client/cli/commands/claude_settings.py b/litellm/proxy/client/cli/commands/claude_settings.py index d0ee507f0b2..e6231f3cac9 100644 --- a/litellm/proxy/client/cli/commands/claude_settings.py +++ b/litellm/proxy/client/cli/commands/claude_settings.py @@ -1,16 +1,18 @@ """Shared handling of Claude Code's ~/.claude/settings.json. `lite up` and `lite autoroute up` patch this file temporarily and restore it on -exit; `lite login --config-claude` and `lite configure claude` patch it -persistently and record how to undo it. All of them need the same merge and the -same apiKeyHelper command, and `up` already imports from `auth`, so the shared -parts live here rather than in any one command module. +exit; `lite configure claude` patches it persistently and records how to undo it. +All of them need the same merge, and `up` already imports from `auth`, so the +shared parts live here rather than in any one command module. The credential is +always a static token in `env.ANTHROPIC_AUTH_TOKEN`: Claude Code's `apiKeyHelper` +would spawn a `lite` process on every credential refresh, and that process touches +the keychain, so nothing here writes one; a helper left by an earlier version is +owned like any other key and stripped. """ import hashlib import json import shlex -import shutil import sys from collections.abc import Callable, Mapping, Sequence from dataclasses import dataclass @@ -27,13 +29,16 @@ from litellm.litellm_core_utils.private_json import ( discard_staged_json, ensure_private_dir, stage_private_json, + write_private_bytes, ) +from . import statusline_script from .cmd_quoting import quote_for_cmd ENV_KEY: Final = "env" API_KEY_HELPER_KEY: Final = "apiKeyHelper" MODEL_KEY: Final = "model" +STATUS_LINE_KEY: Final = "statusLine" ANTHROPIC_BASE_URL_KEY: Final = "ANTHROPIC_BASE_URL" ANTHROPIC_AUTH_TOKEN_KEY: Final = "ANTHROPIC_AUTH_TOKEN" ANTHROPIC_API_KEY_KEY: Final = "ANTHROPIC_API_KEY" @@ -41,6 +46,7 @@ ENABLE_TOOL_SEARCH_KEY: Final = "ENABLE_TOOL_SEARCH" ENABLE_TOOL_SEARCH_VALUE: Final = "true" ENABLE_GATEWAY_MODEL_DISCOVERY_KEY: Final = "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY" ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE: Final = "1" +ANTHROPIC_MODEL_KEY: Final = "ANTHROPIC_MODEL" ANTHROPIC_DEFAULT_MODEL_ENV_KEYS: Final = ( "ANTHROPIC_DEFAULT_SONNET_MODEL", "ANTHROPIC_DEFAULT_HAIKU_MODEL", @@ -53,19 +59,22 @@ OWNED_ENV_KEYS: Final = ( ANTHROPIC_BASE_URL_KEY, ANTHROPIC_AUTH_TOKEN_KEY, ANTHROPIC_API_KEY_KEY, + ANTHROPIC_MODEL_KEY, ) -OWNED_TOP_LEVEL_KEYS: Final = (API_KEY_HELPER_KEY, MODEL_KEY) +OWNED_TOP_LEVEL_KEYS: Final = (API_KEY_HELPER_KEY, MODEL_KEY, STATUS_LINE_KEY) OWNED_PATHS: Final = (*(f"{ENV_KEY}.{key}" for key in OWNED_ENV_KEYS), *OWNED_TOP_LEVEL_KEYS) _CREDENTIAL_ENV_KEYS: Final = frozenset((ANTHROPIC_API_KEY_KEY, ANTHROPIC_AUTH_TOKEN_KEY)) _CREDENTIAL_PATHS: Final = (*(f"{ENV_KEY}.{key}" for key in sorted(_CREDENTIAL_ENV_KEYS)), API_KEY_HELPER_KEY) _BASE_URL_PATH: Final = f"{ENV_KEY}.{ANTHROPIC_BASE_URL_KEY}" -STARTING_MODEL_ROLE: Final = "the /model picker's default row, the model Claude Code starts on" +_MODEL_PATHS: Final = (MODEL_KEY, f"{ENV_KEY}.{ANTHROPIC_MODEL_KEY}") +STARTING_MODEL_ROLE: Final = "the /model picker's default row, the model Claude Code starts and resumes on" CLAUDE_SETTINGS_PATH: Final = Path.home() / ".claude" / "settings.json" CLAUDE_CONFIG_DIR_ENV: Final = "CLAUDE_CONFIG_DIR" BACKUP_PATH: Final = Path.home() / ".litellm" / "claude_settings_backup.json" AUTOROUTE_BACKUP_PATH: Final = Path.home() / ".litellm" / "autorouter" / "claude_settings_backup.json" CONFIGURE_STATE_PATH: Final = Path.home() / ".litellm" / "claude_configure_state.json" +STATUSLINE_SCRIPT_PATH: Final = Path.home() / ".litellm" / "statusline.py" @dataclass(frozen=True, slots=True) @@ -123,16 +132,6 @@ class StaticToken: token: str -@dataclass(frozen=True, slots=True) -class ApiKeyHelper: - """A `lite auth print-token` command Claude Code runs per request, so a login renews in place.""" - - command: str - - -ClaudeCredential: TypeAlias = StaticToken | ApiKeyHelper - - @dataclass(frozen=True, slots=True) class KeepModel: """Leave the top-level `model` as it is, the user's or an earlier configure's (a re-login).""" @@ -145,7 +144,9 @@ class UnpinModel: @dataclass(frozen=True, slots=True) class StartOn: - """Pin the top-level `model`, the row Claude Code starts on.""" + """Pin `model` and `env.ANTHROPIC_MODEL`: the row Claude Code starts on, and the one a resumed session stays + on, since resume otherwise re-sends the transcript's served model, which a raw-model router made a tier + model the key may not reach.""" model: str @@ -259,6 +260,17 @@ def _write_target(settings_path: Path) -> Path: raise ClaudeSettingsError(f"Could not resolve {settings_path}: {e}") from e +def write_claude_settings(settings_path: Path, settings: Mapping[str, JsonValue]) -> None: + """The one way a settings document lands on disk: staged owner-only beside the target and renamed into + place, through a symlink rather than over it. Every writer (`configure`, `up`, `autoroute up` and the + restores) may be carrying the credential, so none creates the file under the umask or truncates it.""" + target: Final = _write_target(settings_path) + try: + commit_staged_json(stage_private_json(str(target), settings), str(target)) + except OSError as e: + raise ClaudeSettingsError(f"Could not write {settings_path}: {e}") from e + + def _stage(path: Path, document: Mapping[str, object]) -> str: try: return stage_private_json(str(path), document) @@ -287,20 +299,49 @@ def _land( raise ClaudeSettingsError(f"Could not {'remove' if staged is None else 'write'} {path}: {e}") from e +def statusline_command(script_path: Path, platform: str = sys.platform) -> str: + """This interpreter, not a bare `python3`: it is the one the apiKeyHelper already depends on.""" + quote: Final = quote_for_cmd if platform.startswith("win") else shlex.quote + return " ".join(quote(token) for token in (sys.executable, str(script_path))) + + +def install_statusline_script(script_path: Path | None = None) -> str: + target: Final = script_path or STATUSLINE_SCRIPT_PATH + try: + ensure_private_dir(target.parent) + write_private_bytes(str(target), Path(statusline_script.__file__).read_bytes()) + except OSError as e: + raise ClaudeSettingsError(f"Could not install the status line script at {target}: {e}") from e + return statusline_command(target) + + +def with_status_line(settings: Mapping[str, JsonValue], command: str) -> Mapping[str, JsonValue]: + """Ours is recognised by the script it runs, so a re-install under another interpreter is still ours.""" + existing: Final = settings.get(STATUS_LINE_KEY) + existing_command: Final = existing.get("command") if isinstance(existing, dict) else None + ours: Final = existing is None or (isinstance(existing_command, str) and command.split()[-1] in existing_command) + if not ours: + return settings + entry: Final = dict((("type", "command"), ("command", command))) # mutable-ok: JSON document + return dict(chain(settings.items(), ((STATUS_LINE_KEY, entry),))) # mutable-ok: JSON document + + def merge_claude_settings( settings: Mapping[str, JsonValue], base_url: str, - credential: ClaudeCredential, + credential: StaticToken, default_model: str | None = None, tier_model: str | None = None, + *, + status_line: str | None = None, ) -> Mapping[str, JsonValue]: """Return a new settings mapping wired to route Claude Code through the proxy. - A StaticToken lands in env.ANTHROPIC_AUTH_TOKEN, an ApiKeyHelper in the top-level apiKeyHelper; - the other credential slots are removed either way, since Claude Code given two credentials may - send the wrong one. ENABLE_TOOL_SEARCH and CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY get their - defaults only when missing. `default_model` is the top-level `model`, the row Claude Code starts - on; `tier_model` is `lite autoroute up`'s knob that points every ANTHROPIC_DEFAULT_*_MODEL at one + The token lands in env.ANTHROPIC_AUTH_TOKEN; the other credential slots (a stray ANTHROPIC_API_KEY, + an apiKeyHelper) are removed, since Claude Code given two credentials may send the wrong one. + ENABLE_TOOL_SEARCH and CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY get their defaults only when + missing. `default_model` is the top-level `model` and env.ANTHROPIC_MODEL (see StartOn); + `tier_model` is `lite autoroute up`'s knob that points every ANTHROPIC_DEFAULT_*_MODEL at one group. Apart from those tier keys, exactly OWNED_PATHS are touched. """ raw_env: Final = settings.get(ENV_KEY, {}) @@ -312,60 +353,24 @@ def merge_claude_settings( (ENABLE_GATEWAY_MODEL_DISCOVERY_KEY, ENABLE_GATEWAY_MODEL_DISCOVERY_VALUE), ), ((key, value) for key, value in current_env.items() if key not in _CREDENTIAL_ENV_KEYS), - ((ANTHROPIC_BASE_URL_KEY, base_url.rstrip("/")),), - ((ANTHROPIC_AUTH_TOKEN_KEY, credential.token),) if isinstance(credential, StaticToken) else (), + ((ANTHROPIC_BASE_URL_KEY, base_url.rstrip("/")), (ANTHROPIC_AUTH_TOKEN_KEY, credential.token)), + ((ANTHROPIC_MODEL_KEY, default_model),) if default_model is not None else (), ((key, tier_model) for key in ANTHROPIC_DEFAULT_MODEL_ENV_KEYS if tier_model is not None), ) ) return dict( # mutable-ok: JSON document handed to json.dump, which rejects a read-only mapping chain( - ((key, value) for key, value in settings.items() if key not in (API_KEY_HELPER_KEY, ENV_KEY)), + ( + (key, value) + for key, value in (with_status_line(settings, status_line) if status_line else settings).items() + if key not in (API_KEY_HELPER_KEY, ENV_KEY) + ), ((ENV_KEY, env),), - ((API_KEY_HELPER_KEY, credential.command),) if isinstance(credential, ApiKeyHelper) else (), ((MODEL_KEY, default_model),) if default_model is not None else (), ) ) -def resolve_api_key_helper(base_url: str, platform: str = sys.platform) -> str: - """Build the shell command Claude Code should run for its apiKeyHelper. - - Claude Code hands the string to the system shell, `sh` on POSIX and cmd.exe - on Windows, so every token is quoted for the shell that will read it. - - Resolves `lite` to an absolute path so the helper works regardless of the - PATH visible to whatever subprocess Claude Code spawns it from. Passing - --base-url explicitly (rather than relying on the bare invocation Claude - Code would otherwise use) makes `print-token` enforce that the cached - token was actually issued for this proxy -- without it, a token minted - for a different, previously-logged-into proxy would be handed to - whichever server the settings currently point at. - - --base-url belongs to the top-level `lite` group, so it has to precede the - subcommand; click rejects it outright after `print-token`. - """ - lite_path: Final = shutil.which("lite") - if lite_path is None: - raise ClaudeSettingsError( - "Could not find `lite` on your PATH. Claude Code's apiKeyHelper needs an absolute path to it." - ) - quote: Final = quote_for_cmd if platform.startswith("win") else shlex.quote - return " ".join(quote(token) for token in (lite_path, "--base-url", base_url, "auth", "print-token")) - - -def lite_api_key_helper_configured(base_url: str, settings_path: Path) -> bool: - """Whether settings_path already carries the apiKeyHelper `lite login --config-claude` writes for base_url. - - Only an exact match counts: a helper for another proxy, a hand-written one, or - settings that cannot be read leave the caller on the env-token path. - """ - try: - configured_helper: Final = load_json_or_empty(settings_path).get(API_KEY_HELPER_KEY) - return configured_helper == resolve_api_key_helper(base_url.rstrip("/")) - except ClaudeSettingsError: - return False - - def _owned(container: Mapping[str, JsonValue], key: str) -> OwnedValue: return OwnedValue(present=key in container, value=container.get(key)) @@ -460,12 +465,13 @@ def read_configure_receipt(state_path: Path) -> ConfigureReceipt | None: def configure_claude_settings( base_url: str, - credential: ClaudeCredential, + credential: StaticToken, model: ModelChoice, settings_path: Path, state_path: Path, owners: Sequence[SettingsFileOwner], commit: Callable[[str, str], None] = commit_staged_json, + script_path: Path | None = None, ) -> None: """Persistently route Claude Code through base_url, recording how to undo it. @@ -474,19 +480,28 @@ def configure_claude_settings( discards the staged settings, and a settings rename that fails after the receipt landed puts the earlier receipt back (or removes the new one), so the receipt on disk never describes settings that were not written. `model`: StartOn pins the starting model, UnpinModel lets go of a pin an - earlier configure made (never of the user's own), KeepModel leaves it alone (a re-login). + earlier configure made (never of the user's own), KeepModel leaves it alone (a re-login). The + status line script is installed and registered under `statusLine` unless the user runs their own; + the receipt owns that key like any other, so unconfigure removes only ours. """ refuse_while_owned(settings_path, owners) current: Final = load_json_or_empty(settings_path) _env_object(current, settings_path) earlier: Final = read_configure_receipt(state_path) - existing: Final = ( - _with(current, MODEL_KEY, earlier.previous[MODEL_KEY]) - if isinstance(model, UnpinModel) and earlier is not None and _ours(current, MODEL_KEY, earlier) - else current + unpinned: Final = MappingProxyType( + { + path: earlier.previous[path] + for path in _MODEL_PATHS + if isinstance(model, UnpinModel) and earlier is not None and _ours(current, path, earlier) + } ) + existing: Final = _with_all(current, unpinned) merged: Final = merge_claude_settings( - existing, base_url, credential, model.model if isinstance(model, StartOn) else None + existing, + base_url, + credential, + model.model if isinstance(model, StartOn) else None, + status_line=install_statusline_script(script_path), ) receipt: Final = _receipt(current, merged, earlier, settings_path.exists()) target: Final = _write_target(settings_path) @@ -581,6 +596,7 @@ __all__ = ( "ANTHROPIC_AUTH_TOKEN_KEY", "ANTHROPIC_BASE_URL_KEY", "ANTHROPIC_DEFAULT_MODEL_ENV_KEYS", + "ANTHROPIC_MODEL_KEY", "API_KEY_HELPER_KEY", "AUTOROUTE_BACKUP_PATH", "BACKUP_PATH", @@ -598,8 +614,8 @@ __all__ = ( "OWNED_TOP_LEVEL_KEYS", "SETTINGS_FILE_OWNERS", "STARTING_MODEL_ROLE", - "ApiKeyHelper", - "ClaudeCredential", + "STATUSLINE_SCRIPT_PATH", + "STATUS_LINE_KEY", "ClaudeSettingsError", "ConfigureReceipt", "KeepModel", @@ -614,12 +630,11 @@ __all__ = ( "claude_settings_path", "configure_claude_settings", "configure_state_path", - "lite_api_key_helper_configured", "load_json_or_empty", "merge_claude_settings", "read_configure_receipt", "refuse_while_owned", - "resolve_api_key_helper", "settings_file_owners", "unconfigure_claude_settings", + "write_claude_settings", ) diff --git a/litellm/proxy/client/cli/commands/configure.py b/litellm/proxy/client/cli/commands/configure.py index 9d329f8d8f2..4acf94e16f9 100644 --- a/litellm/proxy/client/cli/commands/configure.py +++ b/litellm/proxy/client/cli/commands/configure.py @@ -18,11 +18,9 @@ from litellm.proxy.common_utils.model_listing_utils import ( GATEWAY_CLIENT_HEADER, ) -from .auth import CliContextObj, context_secret_vault, get_stored_api_key +from .auth import CliContextObj from .claude_settings import ( STARTING_MODEL_ROLE, - ApiKeyHelper, - ClaudeCredential, ClaudeSettingsError, ModelChoice, StartOn, @@ -33,12 +31,10 @@ from .claude_settings import ( configure_claude_settings, configure_state_path, refuse_while_owned, - resolve_api_key_helper, settings_file_owners, unconfigure_claude_settings, ) from .pi import ListedModel, ListingFailure, PiSyncError, fetch_model_listing -from .up import ensure_fresh_login _LISTED_MODELS_SHOWN: Final = 20 _CLAUDE_TARGET: Final = "claude" @@ -54,24 +50,21 @@ _MODEL_OPTION_HELP: Final = ( ) -def resolve_credential(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, str]: - """The credential to write and the key to check the proxy with. +def resolve_credential(ctx: click.Context, api_key: str | None) -> StaticToken: + """The long-lived key written into settings.json: --api-key, `lite --api-key` or LITELLM_PROXY_API_KEY. - An explicit key (--api-key, `lite --api-key`, LITELLM_PROXY_API_KEY) is long-lived and goes - into settings.json as a static token. Without one, the stored `lite login` credential is used - the way `lite login --config-claude` uses it, through apiKeyHelper, since it expires within a - day and renews in place there; a missing or stale login is refreshed first, as `lite up` does. + A `lite login` credential is never written: it expires within a day, and keeping it fresh would mean + Claude Code running `lite` through `apiKeyHelper` on every credential refresh. """ ctx_obj: Final[CliContextObj] = ctx.obj explicit: Final = api_key or (None if ctx_obj.get("api_key_from_token_file") else ctx_obj.get("api_key")) - if explicit: - return StaticToken(explicit), explicit - base_url: Final = ctx_obj["base_url"] - ensure_fresh_login(ctx) - stored: Final = get_stored_api_key(expected_base_url=base_url, vault=context_secret_vault(ctx)) - if not stored: - raise ClaudeSettingsError("Login did not produce a usable token.") - return ApiKeyHelper(resolve_api_key_helper(base_url)), stored + if not explicit: + raise ClaudeSettingsError( + "`lite configure claude` needs a long-lived virtual key: pass --api-key, `lite --api-key`, or set " + "LITELLM_PROXY_API_KEY. Your `lite login` credential expires within a day, so it is not written " + "into Claude Code's settings." + ) + return StaticToken(explicit) @dataclass(frozen=True, slots=True) @@ -83,22 +76,22 @@ class _Listing: return tuple(model.id for model in self.models) -def _start(ctx: click.Context, api_key: str | None) -> tuple[ClaudeCredential, _Listing]: +def _start(ctx: click.Context, api_key: str | None) -> tuple[StaticToken, _Listing]: """Every configure path begins the same way: the local ownership check first, so a `lite up` - session is refused before any login prompt or request, then the credential, then the listing.""" + session is refused before any request, then the credential, then the listing.""" settings_path: Final = claude_settings_path(os.environ) try: refuse_while_owned(settings_path, settings_file_owners(settings_path)) - credential, key = resolve_credential(ctx, api_key) + credential: Final = resolve_credential(ctx, api_key) except ClaudeSettingsError as e: raise click.ClickException(str(e)) - return credential, _listed_models(ctx.obj["base_url"], key) + return credential, _listed_models(ctx.obj["base_url"], credential.token) def _listing_error(base_url: str, error: PiSyncError) -> str: """The hint that fits how the listing failed: only an unreachable proxy gets the "is it running" question.""" if error.kind is ListingFailure.REJECTED: - return f"LiteLLM rejected your key (HTTP {error.status}). Run `lite login` to refresh it, or pass a valid --api-key." + return f"LiteLLM rejected your key (HTTP {error.status}). Pass a valid --api-key." if error.kind is ListingFailure.UNREACHABLE: return f"{error.message} Is the proxy at {base_url} running, and is --base-url (or LITELLM_PROXY_URL) correct?" if error.kind is ListingFailure.EMPTY: @@ -122,7 +115,7 @@ def _model_choice(model: str | None) -> ModelChoice: return StartOn(model) if model is not None else UnpinModel() -def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listing: _Listing, model: str | None) -> None: +def _apply_claude(ctx: click.Context, credential: StaticToken, listing: _Listing, model: str | None) -> None: ctx_obj: Final[CliContextObj] = ctx.obj base_url: Final = ctx_obj["base_url"] listed: Final = listing.ids @@ -148,16 +141,13 @@ def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listing: _Li in_picker: Final = sum(1 for listed_model in listed if CLAUDE_CODE_PICKER_PATTERN.search(listed_model)) click.echo(f"Configured Claude Code: {settings_path} now routes through {base_url}.") - click.echo( - "Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN." - if isinstance(credential, StaticToken) - else "Credential: your `lite login`, read through apiKeyHelper on every request, so a later login renews it." - ) + click.echo("Credential: your virtual key, stored in the file as ANTHROPIC_AUTH_TOKEN.") click.echo( f"Starting model: {starting} ({STARTING_MODEL_ROLE}); switch any time with /model." if starting is not None else "Starting model: not pinned (Claude Code's default, or a model you set yourself); switch with /model, or " - "pass --model to start on a proxy model." + "pass --model to start on a proxy model. Without a pin, a resumed session re-sends the model its transcript " + "recorded, which behind a raw-model auto-router is the tier model." ) click.echo( f"/model will list all {len(listed)} of the proxy's models." @@ -166,7 +156,7 @@ def _apply_claude(ctx: click.Context, credential: ClaudeCredential, listing: _Li "'claude' or 'anthropic', and this proxy does not list the rest under such names." ) click.echo("Start `claude` from any terminal. Undo with `lite unconfigure claude`.") - if isinstance(credential, StaticToken) and settings_path.is_symlink(): + if settings_path.is_symlink(): click.echo( f"Note: {settings_path} is a symlink to {settings_path.resolve()}, so your key now lives in " "that file; keep it out of version control.", @@ -235,16 +225,16 @@ def unconfigure_group() -> None: "api_key", default=None, help="Long-lived LiteLLM virtual key written into Claude Code's settings. Defaults to the `lite --api-key` / " - "LITELLM_PROXY_API_KEY value; with neither, your `lite login` credential is used through apiKeyHelper.", + "LITELLM_PROXY_API_KEY value; required, since a `lite login` credential expires within a day.", ) @click.option("--model", default=None, help=_MODEL_OPTION_HELP) @click.pass_context def configure_claude(ctx: click.Context, api_key: str | None, model: str | None) -> None: """Route every Claude Code session through your LiteLLM proxy until `lite unconfigure claude`. - Patches ~/.claude/settings.json in place: the proxy URL, your credential (a virtual key as a - static token, or your `lite login` through apiKeyHelper), and gateway model discovery so - /model lists the proxy's models; --model picks the one Claude Code starts on. Every other + Patches ~/.claude/settings.json in place: the proxy URL, your virtual key as a static token, + and gateway model discovery so /model lists the proxy's models; --model picks the one Claude + Code starts on and resumes with. Every other setting is kept, and what changed is recorded so `lite unconfigure claude` can put it back. Assumes the proxy is already running. """ diff --git a/litellm/proxy/client/cli/commands/statusline_script.py b/litellm/proxy/client/cli/commands/statusline_script.py new file mode 100644 index 00000000000..a8abeb68978 --- /dev/null +++ b/litellm/proxy/client/cli/commands/statusline_script.py @@ -0,0 +1,392 @@ +"""Claude Code status line and Codex Stop hook for auto-routed sessions. + +`lite` copies this file verbatim to ~/.litellm/statusline.py and registers it as Claude +Code's `statusLine` command and as Codex's `[[hooks.Stop]]` command, so it must stay +standard-library only and must never import litellm. Claude Code re-runs it on every +status refresh (about every 300ms while typing), so the proxy is asked at most once per +TTL per session and every other refresh is served from a small on-disk cache that holds +only the proxy's answer, never the key. + +Claude Code pipes a JSON payload on stdin (session_id, transcript_path, model); the routed +model is the `message.model` of the latest foreground assistant line in the transcript, +which is the proxy's response `model` field. That only names the tier model when the +auto-router deployment sets `return_raw_model_name: true`; otherwise it is the alias the +client requested. Codex pipes its Stop event instead (hook_event_name, session_id) and has +no transcript to read, so the routed model comes from the proxy's session record and the +result is printed as a `systemMessage` for the transcript. The proxy key is read from the +agent's own environment (the static token `lite configure claude` writes); nothing here +spawns a credential helper. + +Cost figures come from GET /auto_router/session on the proxy, which reads the per-session +rollup written by the spend flush. That flush is asynchronous, so a turn's cost lands a +second or two after the turn; the cache TTL absorbs it. +""" + +from __future__ import annotations + +import hashlib +import json +import os +import sys +import tempfile +import time +import urllib.error +import urllib.request +from collections.abc import Callable, Mapping +from pathlib import Path +from types import MappingProxyType +from typing import IO, Final, NamedTuple, Protocol +from urllib.parse import urlencode + +SESSION_ENDPOINT: Final = "/auto_router/session" +CACHE_TTL_SECONDS: Final = 5.0 +FETCH_TIMEOUT_SECONDS: Final = 3 +BAR_WIDTH: Final = 24 +BAR_FULL: Final = "\u2588" +BAR_EMPTY: Final = "\u2591" +SEPARATOR: Final = " \u00b7 " +TRANSCRIPT_SCAN_LIMIT_BYTES: Final = 4 * 1024 * 1024 +CLAUDE_BASE_URL_ENV_KEYS: Final = ("ANTHROPIC_BASE_URL",) +CLAUDE_API_KEY_ENV_KEYS: Final = ("ANTHROPIC_AUTH_TOKEN", "ANTHROPIC_API_KEY") +CODEX_BASE_URL_ENV_KEYS: Final = ("OPENAI_BASE_URL",) +CODEX_API_KEY_ENV_KEYS: Final = ("OPENAI_API_KEY",) +CODEX_STOP_EVENT: Final = "Stop" +SYNTHETIC_MODEL: Final = "" +LITELLM_LABEL: Final = "LiteLLM" +RESET: Final = "\033[0m" +BOLD: Final = "\033[1m" +DIM: Final = "\033[90m" +LITELLM_COLOR: Final = "\033[38;2;79;70;229m" +BASELINE_COLOR: Final = "\033[38;2;217;119;87m" +EMPTY: Final[Mapping[str, object]] = MappingProxyType({}) +EMPTY_ENV: Final[Mapping[str, str]] = MappingProxyType({}) + + +class Session(NamedTuple): + router_name: str + last_model: str + spend: float + baseline_spend: float + baseline_model: str | None + + +class Credentials(NamedTuple): + base_url: str + api_key: str + + @property + def usable(self) -> bool: + return bool(self.base_url and self.api_key) + + +class Fetched(NamedTuple): + session: Session | None + definitive: bool + + +class Fetch(Protocol): + def __call__(self, credentials: Credentials, session_id: str) -> Fetched: ... + + +def as_mapping(value: object) -> Mapping[str, object]: + return value if isinstance(value, dict) else EMPTY + + +def as_str(value: object) -> str: + return value if isinstance(value, str) else "" + + +def printable(value: object) -> str: + """Labels come from the transcript, the proxy, and Claude Code's model cache, none of which this script + controls, and every one is written to a terminal: a control character (ESC, BEL, C1) in a model name + could redraw the screen or set the clipboard, so only printable text survives.""" + return "".join(character for character in as_str(value) if character.isprintable()) + + +def load_json(raw: bytes | str) -> object: + try: + return json.loads(raw) + except ValueError: + return None + + +def resolve_base_url(env: Mapping[str, str], keys: tuple[str, ...]) -> str: + raw: Final = next((env[key] for key in keys if env.get(key)), "").strip().rstrip("/") + return raw.removesuffix("/v1") + + +def resolve_api_key(env: Mapping[str, str], keys: tuple[str, ...]) -> str: + return next((env[key] for key in keys if env.get(key)), "").strip() + + +def claude_credentials(env: Mapping[str, str]) -> Credentials: + """Claude Code's own resolution order, so the key-scoped lookup runs as the principal that wrote the rows: + ANTHROPIC_AUTH_TOKEN, then ANTHROPIC_API_KEY. A `lite` variable such as LITELLM_PROXY_API_KEY is not a + key Claude Code ever sends, so honoring it would ask as someone else. An apiKeyHelper is never run: a + status line refreshes every few hundred milliseconds, and spawning a credential helper that often is + how a keychain prompt ends up on screen a hundred times.""" + return Credentials(resolve_base_url(env, CLAUDE_BASE_URL_ENV_KEYS), resolve_api_key(env, CLAUDE_API_KEY_ENV_KEYS)) + + +def codex_credentials(env: Mapping[str, str]) -> Credentials: + return Credentials(resolve_base_url(env, CODEX_BASE_URL_ENV_KEYS), resolve_api_key(env, CODEX_API_KEY_ENV_KEYS)) + + +def _transcript_line_model(line: bytes) -> str: + """A `` model is Claude Code's own marker for a locally produced message (an API error, a + resume note), not a served model, so it is skipped like a sidechain line.""" + item: Final = as_mapping(load_json(line)) + if item.get("type") != "assistant" or item.get("isSidechain") is True or item.get("agentId"): + return "" + model: Final = printable(as_mapping(item.get("message")).get("model")) + return "" if model == SYNTHETIC_MODEL else model + + +def latest_transcript_model(transcript_path: str) -> str: + if not transcript_path: + return "" + try: + with Path(transcript_path).open("rb") as transcript: + size: Final = transcript.seek(0, os.SEEK_END) + transcript.seek(max(0, size - TRANSCRIPT_SCAN_LIMIT_BYTES)) + tail: Final = transcript.read() + except OSError: + return "" + return next((model for line in reversed(tail.split(b"\n")) if (model := _transcript_line_model(line))), "") + + +def model_label(model: str, config_dir: Path) -> str: + bare: Final = model.rsplit("/", 1)[-1] + try: + raw: Final = (config_dir / "cache" / "gateway-models.json").read_bytes() + except OSError: + return bare + listed: Final = as_mapping(load_json(raw)).get("models") + if not isinstance(listed, list): + return bare + entries: Final = tuple(as_mapping(entry) for entry in listed) + return next( + ( + printable(entry.get("display_name")) + for entry in entries + if entry.get("id") in (model, bare) and printable(entry.get("display_name")) + ), + bare, + ) + + +def baseline_label(model: str, config_dir: Path) -> str: + labelled: Final = model_label(model, config_dir) + if labelled != model.rsplit("/", 1)[-1]: + return labelled + return " ".join(word.capitalize() for word in labelled.replace("-", " ").split()) + + +def fetch_session(credentials: Credentials, session_id: str) -> Fetched: + """Any 4xx is this credential's definite answer (no row, no access, expired login) and is cached for the + TTL; a 5xx or transport failure is not, so the next refresh tries again.""" + query: Final = urlencode((("session_id", session_id),)) + request: Final = urllib.request.Request( + f"{credentials.base_url}{SESSION_ENDPOINT}?{query}", + headers={ # mutable-ok: urllib.request.Request takes a dict + "Authorization": f"Bearer {credentials.api_key}", + "Accept": "application/json", + }, + ) + try: + with urllib.request.urlopen(request, timeout=FETCH_TIMEOUT_SECONDS) as response: + raw: Final[bytes] = response.read() + except urllib.error.HTTPError as error: + return Fetched(session=None, definitive=400 <= error.code < 500) + except (urllib.error.URLError, OSError): + return Fetched(session=None, definitive=False) + session: Final = _session_from_payload(as_mapping(load_json(raw))) + return Fetched(session=session, definitive=session is not None) + + +def _session_from_payload(payload: Mapping[str, object]) -> Session | None: + router_name: Final = printable(payload.get("router_name")) + last_model: Final = printable(payload.get("last_model")) + spend: Final = payload.get("spend") + baseline_spend: Final = payload.get("baseline_spend") + if not router_name or not last_model: + return None + if not isinstance(spend, (int, float)) or not isinstance(baseline_spend, (int, float)): + return None + return Session( + router_name=router_name, + last_model=last_model, + spend=float(spend), + baseline_spend=float(baseline_spend), + baseline_model=printable(payload.get("baseline_model")) or None, + ) + + +def cache_path(cache_dir: Path, credentials: Credentials, session_id: str) -> Path: + identity: Final = "\n".join((credentials.base_url, credentials.api_key, session_id)) + return cache_dir / hashlib.sha256(identity.encode()).hexdigest() + + +def load_session( + credentials: Credentials, + session_id: str, + cache_dir: Path, + fetch: Fetch = fetch_session, + now: Callable[[], float] = time.time, +) -> Session | None: + path: Final = cache_path(cache_dir, credentials, session_id) + cached: Final = _read_cache(path) + fetched_at: Final = cached.get("fetched_at") + if isinstance(fetched_at, (int, float)) and now() - fetched_at < CACHE_TTL_SECONDS: + return _session_from_payload(as_mapping(cached.get("session"))) + fetched: Final = fetch(credentials, session_id) + if fetched.definitive: + _write_cache(path, fetched.session, now()) + return fetched.session + + +NOFOLLOW: Final = getattr(os, "O_NOFOLLOW", 0) + + +def cache_dir_name() -> str: + return f"litellm-statusline-{os.getuid()}" if hasattr(os, "getuid") else "litellm-statusline" + + +def _own_private_dir(directory: Path) -> bool: + """A shared temp root lets another local user pre-create the directory, so it must be ours and private + before anything is read or written under it. Windows has no uids or POSIX mode bits and a per-user temp + directory already, so there it only has to exist and not be a link.""" + try: + directory.mkdir(mode=0o700, parents=True, exist_ok=True) + status: Final = directory.lstat() + except OSError: + return False + if not os.path.isdir(directory) or os.path.islink(directory): + return False + if not hasattr(os, "getuid"): + return True + return status.st_uid == os.getuid() and not status.st_mode & 0o077 + + +def _read_cache(path: Path) -> Mapping[str, object]: + if not _own_private_dir(path.parent): + return EMPTY + try: + descriptor: Final = os.open(path, os.O_RDONLY | NOFOLLOW) + with os.fdopen(descriptor, "rb") as handle: + return as_mapping(load_json(handle.read())) + except OSError: + return EMPTY + + +def _write_cache(path: Path, session: Session | None, fetched_at: float) -> None: + """Staged beside the entry and renamed into place, so a refresh reading the entry never sees a torn write.""" + entry: Final = session._asdict() if session else None + body: Final = json.dumps({"fetched_at": fetched_at, "session": entry}) # mutable-ok: json.dumps takes a dict + if not _own_private_dir(path.parent): + return + try: + descriptor, staged = tempfile.mkstemp(dir=path.parent, prefix=".tmp-") + except OSError: + return + try: + with os.fdopen(descriptor, "w") as handle: + handle.write(body) + os.replace(staged, path) + except OSError: + Path(staged).unlink(missing_ok=True) + + +def _bar(fraction: float, color: str, width: int, use_color: bool) -> str: + filled: Final = round(max(0.0, min(1.0, fraction)) * width) + if not use_color: + return BAR_FULL * filled + BAR_EMPTY * (width - filled) + return f"{color}{BAR_FULL * filled}{DIM}{BAR_EMPTY * (width - filled)}{RESET}" + + +def render(model: str, session: Session | None, config_dir: Path, use_color: bool, bar_width: int = BAR_WIDTH) -> str: + def paint(code: str, text: str) -> str: + return f"{code}{text}{RESET}" if use_color else text + + routed: Final = paint(BOLD, f"Routed to: {model}") + if session is None: + return routed + header: Final = f"{session.router_name}{SEPARATOR}{routed}" + if session.baseline_model is None or session.baseline_spend <= 0: + return header + reference: Final = baseline_label(session.baseline_model, config_dir) + pct: Final = (session.baseline_spend - session.spend) / session.baseline_spend * 100 + delta: Final = paint(LITELLM_COLOR, f"{'-' if pct >= 0 else '+'}{abs(round(pct))}% vs {reference}") + peak: Final = max(session.spend, session.baseline_spend) + label_width: Final = max(len(LITELLM_LABEL), len(reference)) + rows: Final = ( + (LITELLM_LABEL, session.spend, LITELLM_COLOR), + (reference, session.baseline_spend, BASELINE_COLOR), + ) + lines: Final = ( + f"{paint(DIM, label.ljust(label_width))} {_bar(amount / peak, color, bar_width, use_color)} " + f"{paint(DIM, f'${amount:.2f}')}" + for label, amount, color in rows + ) + return "\n".join((f"{header} {delta}", *lines)) + + +def color_enabled(env: Mapping[str, str]) -> bool: + return env.get("NO_COLOR") is None and env.get("TERM", "") not in ("", "dumb") + + +def status_line( + payload: Mapping[str, object], env: Mapping[str, str], config_dir: Path, cache_dir: Path, fetch: Fetch +) -> str: + fallback: Final = printable(as_mapping(payload.get("model")).get("display_name")) + served: Final = latest_transcript_model(as_str(payload.get("transcript_path"))) + if not served: + return fallback or "claude" + label: Final = model_label(served, config_dir) + session_id: Final = as_str(payload.get("session_id")) + credentials: Final = claude_credentials(env) + if not session_id or not credentials.usable: + return render(label, None, config_dir, color_enabled(env)) + session: Final = load_session(credentials, session_id, cache_dir, fetch) + return render(label, session, config_dir, color_enabled(env)) + + +def codex_stop_message( + payload: Mapping[str, object], env: Mapping[str, str], config_dir: Path, cache_dir: Path, fetch: Fetch +) -> str: + """No cache here: the Stop hook runs once per turn, and a first turn's cached absence would hide the record + the next turn finds.""" + session_id: Final = as_str(payload.get("session_id")) + credentials: Final = codex_credentials(env) + if not session_id or not credentials.usable: + return "" + session: Final = fetch(credentials, session_id).session + if session is None: + return "" + text: Final = render(model_label(session.last_model, config_dir), session, config_dir, use_color=False) + return json.dumps({"systemMessage": f"\n{text}"}) # mutable-ok: json.dumps takes a dict + + +def run(stdin: IO[str], stdout: IO[str], env: Mapping[str, str], fetch: Fetch = fetch_session) -> None: + """A failure renders each mode's own quiet fallback: Claude Code gets the label it already knows, Codex gets + nothing at all rather than a bare string it would reject as hook JSON.""" + body: Final = as_mapping(load_json(stdin.read())) + codex: Final = body.get("hook_event_name") == CODEX_STOP_EVENT + config_dir: Final = Path(env.get("CLAUDE_CONFIG_DIR") or Path.home() / ".claude") + cache_dir: Final = ( + Path(env.get("TMPDIR") or env.get("TEMP") or env.get("TMP") or tempfile.gettempdir()) / cache_dir_name() + ) + try: + text: Final = ( + codex_stop_message(body, env, config_dir, cache_dir, fetch) + if codex + else status_line(body, env, config_dir, cache_dir, fetch) + ) + except Exception: # noqa: BLE001 # a status line must never break the agent session + stdout.write("" if codex else printable(as_mapping(body.get("model")).get("display_name")) or "claude") + return + stdout.write(text) + + +if __name__ == "__main__": + run(sys.stdin, sys.stdout, os.environ) diff --git a/litellm/proxy/client/cli/commands/up.py b/litellm/proxy/client/cli/commands/up.py index ffece87ab83..f2624797a5f 100644 --- a/litellm/proxy/client/cli/commands/up.py +++ b/litellm/proxy/client/cli/commands/up.py @@ -23,11 +23,12 @@ from .auth import CliContextObj, context_secret_vault, get_stored_api_key, load_ from .claude_settings import ( BACKUP_PATH, CLAUDE_SETTINGS_PATH, - ApiKeyHelper, ClaudeSettingsError, + StaticToken, + install_statusline_script, load_json_or_empty, merge_claude_settings, - resolve_api_key_helper, + write_claude_settings, ) @@ -98,8 +99,7 @@ def restore_claude_settings(settings_path: Path | None = None, backup_path: Path return None if record.existed and record.content is not None: resolved_settings_path.parent.mkdir(parents=True, exist_ok=True) - with open(resolved_settings_path, "w") as f: - json.dump(record.content, f, indent=2) + write_claude_settings(resolved_settings_path, record.content) elif resolved_settings_path.exists(): resolved_settings_path.unlink() resolved_backup_path.unlink() @@ -134,13 +134,10 @@ def ensure_fresh_login(ctx: click.Context) -> None: pkce: Final = _stored_login_is_pkce(vault) login_command: Final = "lite login --pkce" if pkce else "lite login" if not sys.stdin.isatty(): - raise UpError( - f"No fresh LiteLLM login found for this proxy. Run `{login_command}` first (apiKeyHelper " - "reads this token on every Claude Code request)." - ) + raise UpError(f"No fresh LiteLLM login found for this proxy. Run `{login_command}` first.") click.echo("No fresh LiteLLM login found for this proxy; starting login...") - ctx.invoke(login, pkce=pkce) + ctx.invoke(login, config_claude=False, pkce=pkce) if not _usable_login(get_stored_api_key(expected_base_url=base_url, vault=vault), vault): raise UpError("Login did not produce a usable token.") @@ -162,7 +159,9 @@ def up(ctx: click.Context) -> None: """Route every Claude Code session through your LiteLLM proxy until stopped. Patches ~/.claude/settings.json so Claude Code picks up the proxy on its own - next startup, from any terminal -- no need to launch it through `lite`. + next startup, from any terminal -- no need to launch it through `lite`. The + key written is the one this command resolved (your fresh `lite login`, or an + explicit --api-key), copied in as a static token for as long as `up` runs. Press Ctrl-C to stop and restore your original settings. Assumes the proxy is already running (this does not start one for you). Cursor is not supported: it has no equivalent file-based config to patch. @@ -180,7 +179,7 @@ def up(ctx: click.Context) -> None: "running (or crashed without cleanup). Run `lite down` first." ) - api_key_helper: Final = resolve_api_key_helper(base_url) + status_line: Final = install_statusline_script() original_existed: Final = CLAUDE_SETTINGS_PATH.exists() original_settings: Final = load_json_or_empty(CLAUDE_SETTINGS_PATH) write_backup( @@ -191,9 +190,10 @@ def up(ctx: click.Context) -> None: ) CLAUDE_SETTINGS_PATH.parent.mkdir(exist_ok=True) - merged: Final = merge_claude_settings(original_settings, base_url, ApiKeyHelper(api_key_helper)) - with open(CLAUDE_SETTINGS_PATH, "w") as f: - json.dump(merged, f, indent=2) + merged: Final = merge_claude_settings( + original_settings, base_url, StaticToken(api_key), status_line=status_line + ) + write_claude_settings(CLAUDE_SETTINGS_PATH, merged) except (AgentRunError, ClaudeSettingsError) as e: raise click.ClickException(str(e)) @@ -248,7 +248,6 @@ __all__ = [ "load_json_or_empty", "merge_claude_settings", "read_backup", - "resolve_api_key_helper", "restore_claude_settings", "up", "write_backup", diff --git a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py index fb2ca6372c0..dbe11882b3c 100644 --- a/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py +++ b/litellm/proxy/common_utils/auth_cache_invalidation_pubsub.py @@ -196,9 +196,7 @@ class AuthCacheInvalidationSubscriber: for additional_cache in self._additional_in_memory_caches: additional_cache.set_cache(parsed.cache_key, parsed.new_value, ttl=parsed.ttl) return - in_memory_cache: Final = self._user_api_key_cache.in_memory_cache - if in_memory_cache is not None: - in_memory_cache.delete_cache(parsed.cache_key) + self._user_api_key_cache.in_memory_cache_for(parsed.cache_key).delete_cache(parsed.cache_key) for additional_cache in self._additional_in_memory_caches: additional_cache.delete_cache(parsed.cache_key) diff --git a/litellm/proxy/common_utils/debug_utils.py b/litellm/proxy/common_utils/debug_utils.py index 554a6ae8d1a..1d4024bb84e 100644 --- a/litellm/proxy/common_utils/debug_utils.py +++ b/litellm/proxy/common_utils/debug_utils.py @@ -147,8 +147,11 @@ async def memory_usage_in_mem_cache( llm_router.cache.in_memory_cache.ttl_dict ) - num_items_in_user_api_key_cache: Final = len(user_api_key_cache.in_memory_cache.cache_dict) + len( - user_api_key_cache.in_memory_cache.ttl_dict + num_items_in_user_api_key_cache: Final = ( + len(user_api_key_cache.in_memory_cache.cache_dict) + + len(user_api_key_cache.in_memory_cache.ttl_dict) + + len(user_api_key_cache.key_object_cache.in_memory_cache.cache_dict) + + len(user_api_key_cache.key_object_cache.in_memory_cache.ttl_dict) ) num_items_in_proxy_logging_obj_cache: Final = len( @@ -189,6 +192,8 @@ async def memory_usage_in_mem_cache_items( return { "user_api_key_cache": user_api_key_cache.in_memory_cache.cache_dict, "user_api_key_ttl": user_api_key_cache.in_memory_cache.ttl_dict, + "user_key_object_cache": user_api_key_cache.key_object_cache.in_memory_cache.cache_dict, + "user_key_object_ttl": user_api_key_cache.key_object_cache.in_memory_cache.ttl_dict, "llm_router_cache": llm_router_in_memory_cache_dict, "llm_router_ttl": llm_router_in_memory_ttl_dict, "proxy_logging_obj_cache": proxy_logging_obj.internal_usage_cache.dual_cache.in_memory_cache.cache_dict, @@ -294,7 +299,9 @@ async def get_memory_summary( try: # User API key cache - user_cache_items: Final = len(user_api_key_cache.in_memory_cache.cache_dict) + user_cache_items: Final = len(user_api_key_cache.in_memory_cache.cache_dict) + len( + user_api_key_cache.key_object_cache.in_memory_cache.cache_dict + ) total_cache_items += user_cache_items caches["user_api_keys"] = { "count": user_cache_items, @@ -429,10 +436,16 @@ def _get_cache_memory_stats( cache_stats: Final[dict[str, object]] = {} try: # User API key cache - user_cache_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.cache_dict) - user_ttl_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.ttl_dict) + key_object_in_memory_cache: Final = user_api_key_cache.key_object_cache.in_memory_cache + user_cache_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.cache_dict) + sys.getsizeof( + key_object_in_memory_cache.cache_dict + ) + user_ttl_size: Final = sys.getsizeof(user_api_key_cache.in_memory_cache.ttl_dict) + sys.getsizeof( + key_object_in_memory_cache.ttl_dict + ) cache_stats["user_api_key_cache"] = { - "num_items": len(user_api_key_cache.in_memory_cache.cache_dict), + "num_items": len(user_api_key_cache.in_memory_cache.cache_dict) + + len(key_object_in_memory_cache.cache_dict), "cache_dict_size_bytes": user_cache_size, "ttl_dict_size_bytes": user_ttl_size, "total_size_mb": round((user_cache_size + user_ttl_size) / (1024 * 1024), 2), diff --git a/litellm/proxy/common_utils/user_api_key_cache.py b/litellm/proxy/common_utils/user_api_key_cache.py index 76982d30306..cb72088ee4a 100644 --- a/litellm/proxy/common_utils/user_api_key_cache.py +++ b/litellm/proxy/common_utils/user_api_key_cache.py @@ -1,11 +1,15 @@ from __future__ import annotations +import re +from collections.abc import Sequence from typing import TYPE_CHECKING, Any, Final, TypeVar, cast, overload from pydantic import BaseModel from litellm._logging import verbose_proxy_logger from litellm.caching.dual_cache import DualCache +from litellm.caching.in_memory_cache import InMemoryCache +from litellm.caching.redis_cache import RedisCache from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL from litellm.proxy.common_utils.cache_pydantic_utils import CacheCodec @@ -14,6 +18,13 @@ if TYPE_CHECKING: T = TypeVar("T", bound=BaseModel) +_HASHED_TOKEN_CACHE_KEY: Final = re.compile(r"[0-9a-f]{64}") + + +def is_user_key_cache_key(key: str) -> bool: + """Only user-key objects are cached under a bare ``hash_token`` digest; every other object uses a prefixed key.""" + return _HASHED_TOKEN_CACHE_KEY.fullmatch(key) is not None + class UserApiKeyCache(DualCache): """ @@ -36,10 +47,50 @@ class UserApiKeyCache(DualCache): ``async_set_cache_pipeline`` applies the same untyped Codec pass as omitting ``model_type`` on ``async_set_cache`` (so ``BaseModel`` rows are dumped before Redis). + User-key objects (see ``is_user_key_cache_key``) live in their own in-memory partition, + ``key_object_cache``, so churn in the other management objects cannot evict them. Both + partitions share the same Redis backend and TTL settings. + ``get_cache`` / ``async_get_cache`` overloads and implementations must be contiguous (no other methods in between) so mypy resolves ``@overload`` + implementation correctly. """ + def __init__( + self, + in_memory_cache: InMemoryCache | None = None, + redis_cache: RedisCache | None = None, + default_in_memory_ttl: float | None = None, + default_redis_ttl: float | None = None, + key_object_in_memory_cache: InMemoryCache | None = None, + ) -> None: + super().__init__( + in_memory_cache=in_memory_cache, + redis_cache=redis_cache, + default_in_memory_ttl=default_in_memory_ttl, + default_redis_ttl=default_redis_ttl, + ) + self.key_object_cache: Final = DualCache( + in_memory_cache=key_object_in_memory_cache or InMemoryCache(), + redis_cache=redis_cache, + default_in_memory_ttl=default_in_memory_ttl, + default_redis_ttl=default_redis_ttl, + ) + + def in_memory_cache_for(self, key: str) -> InMemoryCache: + return self.key_object_cache.in_memory_cache if is_user_key_cache_key(key) else self.in_memory_cache + + def update_cache_ttl(self, default_in_memory_ttl: float | None, default_redis_ttl: float | None) -> None: + super().update_cache_ttl(default_in_memory_ttl=default_in_memory_ttl, default_redis_ttl=default_redis_ttl) + self.key_object_cache.update_cache_ttl( + default_in_memory_ttl=default_in_memory_ttl, default_redis_ttl=default_redis_ttl + ) + + def attach_redis_cache( + self, redis_cache: RedisCache | None = None, *, default_redis_ttl: float | None = None + ) -> None: + super().attach_redis_cache(redis_cache, default_redis_ttl=default_redis_ttl) + self.key_object_cache.attach_redis_cache(redis_cache, default_redis_ttl=default_redis_ttl) + @overload def get_cache( self, @@ -71,7 +122,11 @@ class UserApiKeyCache(DualCache): ) -> object: if model_type is None and "model_type" in kwargs: model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None)) - cached: Final = super().get_cache(key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs) + cached: Final = ( + self.key_object_cache.get_cache(key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs) + if is_user_key_cache_key(key) + else super().get_cache(key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs) + ) if model_type is None: return cached if cached is None: @@ -117,8 +172,14 @@ class UserApiKeyCache(DualCache): ) -> object: if model_type is None and "model_type" in kwargs: model_type = cast(type[BaseModel] | None, kwargs.pop("model_type", None)) - cached: Final = await super().async_get_cache( - key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs + cached: Final = ( + await self.key_object_cache.async_get_cache( + key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs + ) + if is_user_key_cache_key(key) + else await super().async_get_cache( + key=key, parent_otel_span=parent_otel_span, local_only=local_only, **kwargs + ) ) if model_type is None: return cached @@ -137,20 +198,49 @@ class UserApiKeyCache(DualCache): def set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object): model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None)) payload: Final[object] = CacheCodec.serialize(value, model_type=model_type) + if key is not None and is_user_key_cache_key(key): + return self.key_object_cache.set_cache(key=key, value=payload, local_only=local_only, **kwargs) return super().set_cache(key=key, value=payload, local_only=local_only, **kwargs) async def async_set_cache(self, key: str | None, value: object, local_only: bool = False, **kwargs: object): model_type: Final = cast(type[BaseModel] | None, kwargs.pop("model_type", None)) payload: Final[object] = CacheCodec.serialize(value, model_type=model_type) + if key is not None and is_user_key_cache_key(key): + return await self.key_object_cache.async_set_cache(key=key, value=payload, local_only=local_only, **kwargs) return await super().async_set_cache(key=key, value=payload, local_only=local_only, **kwargs) - async def async_set_cache_pipeline(self, cache_list: list, local_only: bool = False, **kwargs: object) -> None: + def delete_cache(self, key: str) -> None: + if is_user_key_cache_key(key): + self.key_object_cache.delete_cache(key) + return + super().delete_cache(key) + + async def async_delete_cache(self, key: str) -> None: + if is_user_key_cache_key(key): + await self.key_object_cache.async_delete_cache(key) + return + await super().async_delete_cache(key) + + def flush_cache(self) -> None: + super().flush_cache() + self.key_object_cache.in_memory_cache.flush_cache() + + async def async_set_cache_pipeline( + self, cache_list: Sequence[tuple[str, object]], local_only: bool = False, **kwargs: object + ) -> None: """ Batch writes with the same Codec boundary as ``async_set_cache`` without ``model_type``: ``BaseModel`` values become JSON-safe dicts; dicts/scalars unchanged. """ - normalized: Final = [(key, CacheCodec.serialize(value, model_type=None)) for key, value in cache_list] - return await super().async_set_cache_pipeline(cache_list=normalized, local_only=local_only, **kwargs) + normalized: Final = tuple((key, CacheCodec.serialize(value, model_type=None)) for key, value in cache_list) + key_object_entries: Final = tuple(entry for entry in normalized if is_user_key_cache_key(entry[0])) + other_entries: Final = tuple(entry for entry in normalized if not is_user_key_cache_key(entry[0])) + if key_object_entries: + await self.key_object_cache.async_set_cache_pipeline( + cache_list=key_object_entries, local_only=local_only, **kwargs + ) + if other_entries: + await super().async_set_cache_pipeline(cache_list=other_entries, local_only=local_only, **kwargs) #: Value cached under ``user_object_permission_id_cache_key`` when the user links no permission row, diff --git a/litellm/proxy/db/autorouter_session_rollup.py b/litellm/proxy/db/autorouter_session_rollup.py index b866ecc741f..3a61da164d0 100644 --- a/litellm/proxy/db/autorouter_session_rollup.py +++ b/litellm/proxy/db/autorouter_session_rollup.py @@ -103,6 +103,7 @@ class AutoRouterTurnTransaction: cache_ttl_seconds: int | None cache_touched: bool tier: str | None = None + baseline_model: str | None = None class TurnCacheFacts(NamedTuple): @@ -168,7 +169,7 @@ def _write_ttl_seconds(usage_object: Mapping[str, object] | None) -> int | None: SESSION_ID_MAX_CHARS: Final = 256 -def _bounded_session_id(session_id: str) -> str: +def bounded_session_id(session_id: str) -> str: """The session id as stored, bounded so a caller-chosen identifier cannot exceed Postgres's B-tree index entry limit through the composite primary key. Oversized ids map to a stable digest, so their turns still aggregate into one session.""" @@ -194,6 +195,9 @@ def build_autorouter_turn_transaction( classifier_cost folded into this turn's spend: the excluded classifier row is how it was billed, the decision is how it is attributed. Cache facts are derived from the payload's own usage record through the savings owner, never handed in beside it. + The baseline the turn's saved_spend was priced against travels with the turn, so the + row can name the counterfactual for the money it holds even after the router is + reconfigured or removed. """ if payload.get("status") != "success": return None @@ -216,13 +220,15 @@ def build_autorouter_turn_transaction( usage_object_raw: Final = metadata.get("usage_object") cache: Final = turn_cache_facts(usage_object_raw if isinstance(usage_object_raw, Mapping) else None) tier_raw: Final = routing_decision.get("tier") + baseline_raw: Final = routing_decision.get("savings_baseline_model") classifier_cost: Final = classifier_cost_from_decision(routing_decision) return AutoRouterTurnTransaction( api_key=api_key, - session_id=_bounded_session_id(session_id), + session_id=bounded_session_id(session_id), router_name=router_name, router_type=str(routing_decision.get("router_type") or "unknown"), tier=tier_raw if isinstance(tier_raw, str) and tier_raw else None, + baseline_model=baseline_raw if isinstance(baseline_raw, str) and baseline_raw else None, model=model, turn_at=turn_at, total_tokens=int(payload.get("prompt_tokens") or 0) + int(payload.get("completion_tokens") or 0), @@ -253,6 +259,10 @@ _CACHE_TTL: Final = _p("cache_ttl_seconds") _TOUCHED: Final = _p("cache_touched") _TIER: Final = f"{_p('tier')}::text" _TIER_DELTA: Final = f"(CASE WHEN {_TIER} IS NULL THEN '{{}}'::jsonb ELSE jsonb_build_object({_TIER}, 1) END)" +_BASELINE: Final = f"{_p('baseline_model')}::text" +_BASELINE_DELTA: Final = ( + f"(CASE WHEN {_BASELINE} IS NULL THEN '{{}}'::jsonb ELSE jsonb_build_object({_BASELINE}, 1) END)" +) _IN_ORDER: Final = f"{_TURN_AT}::timestamp >= t.last_turn_at" _SAME: Final = f"{_IN_ORDER} AND t.last_model = {_MODEL}" @@ -270,7 +280,8 @@ INSERT INTO "LiteLLM_AutoRouterSession" AS t ( last_model, models, turns, unordered_turns, covered_turns, cache_hits, same_model_turns, same_model_hits, first_visit_turns, first_visit_hits, return_turns, return_hits, return_expired_misses, return_within_ttl_misses, - ttl_5m_turns, ttl_1h_turns, total_tokens, spend, saved_spend, classifier_cost, classifier_cost_recorded_turns, tier_turns + ttl_5m_turns, ttl_1h_turns, total_tokens, spend, saved_spend, classifier_cost, classifier_cost_recorded_turns, tier_turns, + baseline_models ) VALUES ( {_p("api_key")}, {_p("session_id")}, {_p("router_name")}, {_p("router_type")}, {_TURN_AT}::timestamp, {_TURN_AT}::timestamp, @@ -281,7 +292,7 @@ VALUES ( (CASE WHEN {_CACHE_TTL}::int = {CACHE_TTL_5M_SECONDS} THEN 1 ELSE 0 END), (CASE WHEN {_CACHE_TTL}::int = {CACHE_TTL_1H_SECONDS} THEN 1 ELSE 0 END), {_p("total_tokens")}::bigint, {_p("spend")}::float8, {_p("saved_spend")}::float8, - {_p("classifier_cost")}::float8, 1, {_TIER_DELTA} + {_p("classifier_cost")}::float8, 1, {_TIER_DELTA}, {_BASELINE_DELTA} ) ON CONFLICT (api_key, session_id, router_name) DO UPDATE SET turns = t.turns + 1, @@ -317,6 +328,9 @@ ON CONFLICT (api_key, session_id, router_name) DO UPDATE SET tier_turns = (CASE WHEN {_TIER} IS NOT NULL AND t.router_type = {_p("router_type")} THEN t.tier_turns || jsonb_build_object({_TIER}, COALESCE((t.tier_turns ->> {_TIER})::int, 0) + 1) ELSE t.tier_turns END), + baseline_models = (CASE WHEN {_BASELINE} IS NOT NULL + THEN t.baseline_models || jsonb_build_object({_BASELINE}, COALESCE((t.baseline_models ->> {_BASELINE})::int, 0) + 1) + ELSE t.baseline_models END), first_turn_at = LEAST(t.first_turn_at, EXCLUDED.first_turn_at), last_turn_at = GREATEST(t.last_turn_at, EXCLUDED.last_turn_at) """ diff --git a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py index 6e29d44662e..de9618a44a1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py +++ b/litellm/proxy/guardrails/guardrail_hooks/azure/prompt_shield.py @@ -338,10 +338,9 @@ class AzureContentSafetyPromptShieldGuardrail(AzureGuardrailBase, CustomGuardrai estimated cost) and the ``azure`` provider label to the recorded guardrail information. Follows the OpenAI moderation override pattern (openai/moderations.py).""" - guardrail_response: Final[dict | str] = ( # mutable-ok: mirrors CustomGuardrail._process_response - ("mask" if self._inputs_were_modified(original_inputs, response) else "allow") - if original_inputs is not None and isinstance(response, dict) - else ({} if response is None else response) # mutable-ok: empty placeholder, never mutated + guardrail_response: Final = self._summarize_guardrail_response( + response=response, + original_inputs=original_inputs, ) self.add_standard_logging_guardrail_information_to_request_data( guardrail_json_response=guardrail_response, diff --git a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py index 722f96ef814..1e684c514de 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py +++ b/litellm/proxy/guardrails/guardrail_hooks/litellm_content_filter/content_filter.py @@ -9,6 +9,7 @@ import asyncio import json import os import re +import time from collections.abc import AsyncGenerator, Coroutine, Mapping, Sequence from datetime import datetime from re import Pattern @@ -1702,6 +1703,7 @@ class ContentFilterGuardrail(CustomGuardrail): start_time: datetime, masked_entity_count: dict[str, int], exception_str: str, + duration: float | None = None, ) -> None: """ Log guardrail information to request_data metadata. @@ -1713,6 +1715,7 @@ class ContentFilterGuardrail(CustomGuardrail): start_time: Start time of guardrail execution masked_entity_count: Count of masked entities by type exception_str: Exception string if guardrail failed + duration: Seconds spent inside the guardrail; defaults to the wall clock since start_time """ # Convert TypedDict detections to regular dicts for JSON serialization guardrail_json_response: Exception | str | dict | list[dict] = [dict(detection) for detection in detections] @@ -1741,7 +1744,7 @@ class ContentFilterGuardrail(CustomGuardrail): guardrail_status=status, start_time=start_time.timestamp(), end_time=datetime.now().timestamp(), - duration=(datetime.now() - start_time).total_seconds(), + duration=(datetime.now() - start_time).total_seconds() if duration is None else duration, masked_entity_count=masked_entity_count, tracing_detail=GuardrailTracingDetail(**tracing_kw), ) @@ -1971,6 +1974,7 @@ class ContentFilterGuardrail(CustomGuardrail): buffer_size: Final = 50 # Increased buffer to catch patterns split across many chunks start_time: Final = datetime.now() + scan_seconds: float = 0.0 # rebind-ok: accumulates per-chunk scan time across the stream detections: list[ContentFilterDetection] = [] masked_entity_count: Final[dict[str, int]] = {} status: GuardrailStatus = "success" @@ -2007,6 +2011,7 @@ class ContentFilterGuardrail(CustomGuardrail): # Add a space at the end if it's the final chunk to trigger word boundaries (\b) text_to_scan = text_to_check + (" " if is_final else "") choice_detections: list[ContentFilterDetection] = [] + scan_started = time.perf_counter() try: # _filter_single_text scans the whole accumulated @@ -2024,6 +2029,8 @@ class ContentFilterGuardrail(CustomGuardrail): except Exception as e: verbose_proxy_logger.error("ContentFilterGuardrail: Error in masking: %s", e) masked_text = text_to_scan # Fallback to current text + finally: + scan_seconds += time.perf_counter() - scan_started # Determine how much can be safely yielded if is_final: @@ -2074,6 +2081,7 @@ class ContentFilterGuardrail(CustomGuardrail): start_time=start_time, masked_entity_count=masked_entity_count, exception_str=exception_str, + duration=scan_seconds, ) @staticmethod diff --git a/litellm/proxy/hooks/proxy_track_cost_callback.py b/litellm/proxy/hooks/proxy_track_cost_callback.py index 00406ad436e..1e946cc2e23 100644 --- a/litellm/proxy/hooks/proxy_track_cost_callback.py +++ b/litellm/proxy/hooks/proxy_track_cost_callback.py @@ -1,6 +1,6 @@ import asyncio import traceback -from collections.abc import Sequence +from collections.abc import Callable, Sequence from datetime import datetime from typing import TYPE_CHECKING, Any, Final, cast @@ -23,6 +23,7 @@ from litellm.proxy.auth.auth_checks import ( ) from litellm.proxy.auth.route_checks import RouteChecks from litellm.proxy.db.db_spend_update_writer import ( + DBSpendUpdateWriter, debitable_model_access_groups, get_llm_router, ) @@ -81,6 +82,12 @@ _CAPTURED_IDENTITY_CALL_TYPES: Final[frozenset[str]] = frozenset( ) +def _proxy_spend_writer() -> DBSpendUpdateWriter: + from litellm.proxy.proxy_server import proxy_logging_obj + + return proxy_logging_obj.db_spend_update_writer + + class _ProxyDBLogger(CustomLogger): def __init__( self, @@ -88,9 +95,11 @@ class _ProxyDBLogger(CustomLogger): *, turn_off_message_logging: bool = False, message_logging: bool = True, + spend_writer: Callable[[], DBSpendUpdateWriter] = _proxy_spend_writer, ) -> None: super().__init__(turn_off_message_logging=turn_off_message_logging, message_logging=message_logging) self.spend_event_producer = spend_event_producer + self._spend_writer: Final = spend_writer async def async_log_success_event( self, kwargs: ObjectMapping, response_obj: object, start_time: datetime, end_time: datetime @@ -150,8 +159,6 @@ class _ProxyDBLogger(CustomLogger): ): return - from litellm.proxy.proxy_server import proxy_logging_obj - _metadata = dict( LiteLLMProxyRequestSetup.get_sanitized_user_information_from_key(user_api_key_dict=user_api_key_dict) ) @@ -227,13 +234,12 @@ class _ProxyDBLogger(CustomLogger): if request_data.get("litellm_trace_id") is None: request_data["litellm_trace_id"] = getattr(_litellm_logging_obj, "litellm_trace_id", None) - # Use the actual request start time from the logging object so that - # failed requests record the real duration instead of 0. - actual_start_time = datetime.now() - if _litellm_logging_obj is not None: - obj_start: Final = getattr(_litellm_logging_obj, "start_time", None) - if obj_start is not None: - actual_start_time = obj_start + lifted_start_time: Final = request_data.get("start_time") + actual_start_time: Final = ( + lifted_start_time + if isinstance(lifted_start_time, datetime) + else getattr(_litellm_logging_obj, "start_time", None) or datetime.now() + ) # A stream that broke mid-flight still billed the provider for the # chunks already delivered. ``post_call_failure_hook`` lifts that @@ -249,7 +255,7 @@ class _ProxyDBLogger(CustomLogger): existing_metadata.get("standard_logging_guardrail_information") ) - await proxy_logging_obj.db_spend_update_writer.update_database( + await self._spend_writer().update_database( token=user_api_key_dict.api_key, response_cost=recovered_response_cost, user_id=user_api_key_dict.user_id, diff --git a/litellm/proxy/hooks/responses_id_security.py b/litellm/proxy/hooks/responses_id_security.py index c4c15c40d1e..7e7f70d6f7e 100644 --- a/litellm/proxy/hooks/responses_id_security.py +++ b/litellm/proxy/hooks/responses_id_security.py @@ -32,6 +32,29 @@ if TYPE_CHECKING: _RESPONSES_API_PROVIDER_PREFIX: Final = "/openai" _RESPONSES_API_CREATE_ROUTES: Final = frozenset({"/v1/responses", "/responses"}) +_ADDRESSED_RESPONSE_ID_KEY: Final = "_litellm_addressed_response_id" +_UNMANAGED_RESPONSE_ID_DETAIL: Final = ( + "Forbidden. This response id was not issued by this proxy, so the proxy cannot tell who owns it. " + "To let keys address responses this proxy did not issue, set " + "general_settings::allow_unmanaged_response_ids to True in the config.yaml file." +) +_PROXY_ADMIN_ROLES: Final = frozenset({LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN.value}) + + +def _proxy_general_settings() -> Mapping[str, Any]: + from litellm.proxy.proxy_server import general_settings + + return general_settings + + +def _proxy_signing_key() -> str | None: + import os + + from litellm.proxy.proxy_server import master_key + + salt_key: Final = os.getenv("LITELLM_SALT_KEY", None) + return master_key if salt_key is None else salt_key + _RESPONSE_PAYLOAD_ADAPTER: Final = TypeAdapter(Mapping[str, object]) @@ -83,8 +106,13 @@ def _is_responses_api_create_route(request_route: str | None) -> bool: class ResponsesIDSecurity(CustomLogger): - def __init__(self): - pass + def __init__( + self, + general_settings_reader: Callable[[], Mapping[str, Any]] = _proxy_general_settings, + signing_key_reader: Callable[[], str | None] = _proxy_signing_key, + ) -> None: + self._general_settings_reader: Final = general_settings_reader + self._signing_key_reader: Final = signing_key_reader async def async_pre_call_hook( self, @@ -103,30 +131,51 @@ class ResponsesIDSecurity(CustomLogger): } if call_type not in responses_api_call_types: return None - if call_type == "aresponses": - # check 'previous_response_id' if present in the data - previous_response_id: Final = data.get("previous_response_id") - if previous_response_id and self._is_encrypted_response_id(previous_response_id): - original_response_id, user_id, team_id = self._decrypt_response_id(previous_response_id) - self.check_user_access_to_response_id(user_id, team_id, user_api_key_dict) - data["previous_response_id"] = original_response_id - elif call_type in {"aget_responses", "adelete_responses", "acancel_responses", "alist_input_items"}: - response_id: Final = data.get("response_id") - - if response_id and self._is_encrypted_response_id(response_id): - original_response_id, user_id, team_id = self._decrypt_response_id(response_id) - - self.check_user_access_to_response_id(user_id, team_id, user_api_key_dict) - data["response_id"] = original_response_id + addressed_id_field: Final = "previous_response_id" if call_type == "aresponses" else "response_id" + retained_id: Final = data.get(_ADDRESSED_RESPONSE_ID_KEY) + addressed_id: Final = ( + retained_id if isinstance(retained_id, str) and retained_id else data.get(addressed_id_field) + ) + if not isinstance(addressed_id, str) or not addressed_id: + return data + authorized_id: Final = self._authorize_response_id(addressed_id, user_api_key_dict) + data[addressed_id_field] = authorized_id + data[_ADDRESSED_RESPONSE_ID_KEY] = addressed_id return data + def _authorize_response_id( + self, + response_id: str, + user_api_key_dict: "UserAPIKeyAuth", + ) -> str: + if self._is_encrypted_response_id(response_id): + original_response_id, user_id, team_id = self._decrypt_response_id(response_id) + self.check_user_access_to_response_id(user_id, team_id, user_api_key_dict) + return original_response_id + + if self._unmanaged_response_ids_allowed(user_api_key_dict): + return response_id + + raise HTTPException(status_code=403, detail=_UNMANAGED_RESPONSE_ID_DETAIL) + + def _unmanaged_response_ids_allowed(self, user_api_key_dict: "UserAPIKeyAuth") -> bool: + general_settings: Final = self._general_settings_reader() + + if general_settings.get("disable_responses_id_security", False): + return True + if general_settings.get("allow_unmanaged_response_ids", False): + return True + if self._get_signing_key() is None: + return True + return user_api_key_dict.user_role in _PROXY_ADMIN_ROLES + def check_user_access_to_response_id( self, response_id_user_id: str | None, response_id_team_id: str | None, user_api_key_dict: "UserAPIKeyAuth", ) -> bool: - from litellm.proxy.proxy_server import general_settings + general_settings: Final = self._general_settings_reader() if ( user_api_key_dict.user_role == LitellmUserRoles.PROXY_ADMIN.value @@ -219,15 +268,7 @@ class ResponsesIDSecurity(CustomLogger): return response_id, None, None def _get_signing_key(self) -> str | None: - """Get the signing key for encryption/decryption.""" - import os - - from litellm.proxy.proxy_server import master_key - - salt_key = os.getenv("LITELLM_SALT_KEY", None) - if salt_key is None: - salt_key = master_key - return salt_key + return self._signing_key_reader() def _encrypt_response_id( self, @@ -274,7 +315,7 @@ class ResponsesIDSecurity(CustomLogger): This method adds response IDs to an in-memory queue, which are then batch-processed by the DBSpendUpdateWriter during regular database update cycles. """ - from litellm.proxy.proxy_server import general_settings + general_settings: Final = self._general_settings_reader() if general_settings.get("disable_responses_id_security", False): return response @@ -288,7 +329,7 @@ class ResponsesIDSecurity(CustomLogger): async def async_post_call_streaming_iterator_hook( self, user_api_key_dict: "UserAPIKeyAuth", response: Any, request_data: dict ) -> AsyncGenerator[BaseLiteLLMOpenAIResponseObject, None]: - from litellm.proxy.proxy_server import general_settings + general_settings: Final = self._general_settings_reader() # Create a request-scoped cache for consistent encryption across streaming chunks. request_encryption_cache: Final[dict[str, str]] = {} diff --git a/litellm/proxy/litellm_pre_call_utils.py b/litellm/proxy/litellm_pre_call_utils.py index ad4687e95db..8f7b515c22a 100644 --- a/litellm/proxy/litellm_pre_call_utils.py +++ b/litellm/proxy/litellm_pre_call_utils.py @@ -173,6 +173,7 @@ _ENABLE_TEAM_STALE_ALIAS_BYPASS: bool | None = None if TYPE_CHECKING: from litellm.integrations.otel.model.destination import OtelDestination + from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry from litellm.proxy.proxy_server import ProxyConfig as _ProxyConfig from litellm.types.proxy.policy_engine import Policy, PolicyMatchContext @@ -2015,6 +2016,7 @@ async def add_litellm_data_to_request( "method": request.method, "headers": _logging_safe_headers, "body": None, # filled in post-strip; see below + "credential_fields": tuple(sorted(name for name in _TRANSPORT_ONLY_CREDENTIAL_KEYS if name in data)), "arrival_time": arrival_time, # Track when request arrived at proxy } @@ -3141,6 +3143,7 @@ def _match_and_track_policies( context: "PolicyMatchContext", request_body_policies: Sequence[str], policies_override: dict[str, "Policy"] | None = None, + attachment_registry_override: "AttachmentRegistry | None" = None, ) -> tuple[list[str], dict[str, str]]: """ Match policies via attachments and request body, track them in metadata. @@ -3157,7 +3160,9 @@ def _match_and_track_policies( from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher # Get matching policies via attachments (with match reasons for attribution) - attachment_registry: Final = get_attachment_registry() + attachment_registry: Final = ( + attachment_registry_override if attachment_registry_override is not None else get_attachment_registry() + ) matches_with_reasons: Final = attachment_registry.get_attached_policies_with_reasons(context) matching_policy_names: Final = [m["policy_name"] for m in matches_with_reasons] policy_reasons: Final = {m["policy_name"]: m["matched_via"] for m in matches_with_reasons} @@ -3165,9 +3170,11 @@ def _match_and_track_policies( verbose_proxy_logger.debug("Policy engine: matched policies via attachments: %s", matching_policy_names) # Combine attachment-based policies with dynamic request body policies - all_policy_names: Final = set(matching_policy_names) - if request_body_policies and isinstance(request_body_policies, list): - all_policy_names.update(request_body_policies) + request_body_policies_list: Final = ( + tuple(request_body_policies) if request_body_policies and isinstance(request_body_policies, list) else () + ) + all_policy_names: Final = tuple(dict.fromkeys((*matching_policy_names, *request_body_policies_list))) + if request_body_policies_list: verbose_proxy_logger.debug("Policy engine: added dynamic policies from request body: %s", request_body_policies) if not all_policy_names: @@ -3238,16 +3245,14 @@ def _apply_resolved_guardrails_to_metadata( if not resolved_guardrails and not pipelines: return - existing_guardrails = data[metadata_variable_name].get("guardrails", []) - if not isinstance(existing_guardrails, list): - existing_guardrails = [] + existing_guardrails: Final = data[metadata_variable_name].get("guardrails", []) + existing_guardrails_list: Final = existing_guardrails if isinstance(existing_guardrails, list) else [] # Combine existing guardrails with policy-resolved guardrails (no duplicates) - combined = set(existing_guardrails) - combined.update(resolved_guardrails) - data[metadata_variable_name]["guardrails"] = list(combined) + combined: Final = list(dict.fromkeys((*existing_guardrails_list, *resolved_guardrails))) + data[metadata_variable_name]["guardrails"] = combined - verbose_proxy_logger.debug("Policy engine: added guardrails to request metadata: %s", list(combined)) + verbose_proxy_logger.debug("Policy engine: added guardrails to request metadata: %s", combined) async def add_guardrails_from_policy_engine( diff --git a/litellm/proxy/management_endpoints/auto_router_endpoints.py b/litellm/proxy/management_endpoints/auto_router_endpoints.py index bbc914a772a..50716e5d474 100644 --- a/litellm/proxy/management_endpoints/auto_router_endpoints.py +++ b/litellm/proxy/management_endpoints/auto_router_endpoints.py @@ -32,11 +32,15 @@ from litellm.proxy.auth.auth_checks import ( can_key_call_resolved_model, ) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth -from litellm.proxy.db.autorouter_session_rollup import AUTOROUTER_BENCHMARKS_SQL +from litellm.proxy.db.autorouter_session_rollup import ( + AUTOROUTER_BENCHMARKS_SQL, + bounded_session_id, +) from litellm.proxy.litellm_pre_call_utils import ( LiteLLMProxyRequestSetup, refresh_proxy_server_request_body_snapshot, ) +from litellm.repositories.autorouter_session_repository import AutoRouterSessionRepository from litellm.repositories.base_repository import SupportsModelDump from litellm.repositories.team_repository import TeamRepository from litellm.router_strategy.complexity_router import ComplexityRouter @@ -54,6 +58,7 @@ from litellm.types.management_endpoints.auto_router_endpoints import ( AutoRouterCacheStats, AutoRouterRoutingTestRequest, AutoRouterRoutingTestResponse, + AutoRouterSessionResponse, ComplexityRouterConfigValidationRequest, ComplexityRouterConfigValidationResponse, RequestComplexityRouterConfig, @@ -704,6 +709,51 @@ async def get_auto_router_benchmarks( ) +@router.get( + "/auto_router/session", + tags=("auto router",), + dependencies=(Depends(user_api_key_auth),), + response_model=AutoRouterSessionResponse, +) +async def get_auto_router_session( + user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + session_id: Annotated[ + str, Query(description="The client session id (x-*-session-id header) the turns were sent under") + ], +) -> AutoRouterSessionResponse: + """ + One auto-routed session, for the key that ran it: the model its last turn was routed to and the + session's spend against the router's savings baseline. Built for a coding agent's status line + or stop hook, so any virtual key may call it and only ever sees rows written under its own + key hash. Reads the LiteLLM_AutoRouterSession rollup, which the asynchronous spend flush + fills a moment after each turn; a session with no flushed auto-routed turn yet is a 404. The + id is bounded the way the writer bounded it, so an oversized client id still finds its row. + """ + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) + row: Final = await AutoRouterSessionRepository(prisma_client).find_latest_for_key( + user_api_key_dict.api_key, bounded_session_id(session_id) + ) + if row is None: + raise HTTPException( + status_code=404, detail=f"No auto-routed turns recorded for session {session_id!r} under this key" + ) + return AutoRouterSessionResponse( + session_id=session_id, + router_name=row.router_name, + router_type=row.router_type, + turns=row.turns, + last_model=row.last_model, + spend=row.spend, + saved_spend=row.saved_spend, + baseline_spend=row.spend + row.saved_spend, + baseline_model=row.baseline_model, + baseline_models=row.baseline_models, + ) + + # --------------------------------------------------------------------------- # Shadow eval: pre-adoption evaluation of an auto-router against live traffic. # The job row is immutable config plus stopped_at; status, counts, spend, and errors diff --git a/litellm/proxy/management_endpoints/key_management_endpoints.py b/litellm/proxy/management_endpoints/key_management_endpoints.py index 749a940de0e..db4467dda5b 100644 --- a/litellm/proxy/management_endpoints/key_management_endpoints.py +++ b/litellm/proxy/management_endpoints/key_management_endpoints.py @@ -3729,7 +3729,7 @@ async def delete_key_fn( ) verbose_proxy_logger.debug( - "/keys/delete - cache after delete: %s", user_api_key_cache.in_memory_cache.cache_dict + "/keys/delete - cache after delete: %s", user_api_key_cache.key_object_cache.in_memory_cache.cache_dict ) asyncio.create_task( diff --git a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py index face515ef88..28a8bab1f24 100644 --- a/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py +++ b/litellm/proxy/pass_through_endpoints/llm_passthrough_endpoints.py @@ -97,7 +97,6 @@ else: vertex_llm_base: Final = VertexBase() router: Final = APIRouter() -openai_passthrough_router: Final = APIRouter() default_vertex_config: Final = None passthrough_endpoint_router: Final = PassthroughEndpointRouter() @@ -2297,11 +2296,6 @@ async def vertex_proxy_route( ) -@openai_passthrough_router.api_route( - "/openai_passthrough/{endpoint:path}", - methods=["GET", "POST", "PUT", "DELETE", "PATCH"], - tags=["OpenAI Pass-through", "pass-through"], -) @router.api_route( "/openai/{endpoint:path}", methods=["GET", "POST", "PUT", "DELETE", "PATCH"], diff --git a/litellm/proxy/pass_through_endpoints/openai_passthrough_endpoints.py b/litellm/proxy/pass_through_endpoints/openai_passthrough_endpoints.py new file mode 100644 index 00000000000..f56a59dd560 --- /dev/null +++ b/litellm/proxy/pass_through_endpoints/openai_passthrough_endpoints.py @@ -0,0 +1,44 @@ +"""/openai_passthrough must be matched ahead of the native /{provider}/v1/files and +/{provider}/v1/batches routes, so unlike the other provider passthrough routes it is +registered at startup and defers to the lazily loaded handler per call.""" + +from typing import Final + +from fastapi import APIRouter, Depends, Request, Response + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth + +router: Final = APIRouter() + + +@router.api_route( + "/openai_passthrough/{endpoint:path}", + methods=["GET", "POST", "PUT", "DELETE", "PATCH"], + tags=["OpenAI Pass-through", "pass-through"], +) +async def openai_passthrough_route( + endpoint: str, + request: Request, + fastapi_response: Response, + user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), +) -> Response: + """ + Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native + implementations (e.g. the Responses API at /v1/responses). + + Examples: + - /openai_passthrough/v1/responses + - /openai_passthrough/v1/responses/{response_id} + - /openai_passthrough/v1/responses/{response_id}/input_items + + [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough) + """ + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import openai_proxy_route + + return await openai_proxy_route( + endpoint=endpoint, + request=request, + fastapi_response=fastapi_response, + user_api_key_dict=user_api_key_dict, + ) diff --git a/litellm/proxy/policy_engine/attachment_registry.py b/litellm/proxy/policy_engine/attachment_registry.py index 05df44242aa..a8ead86ac36 100644 --- a/litellm/proxy/policy_engine/attachment_registry.py +++ b/litellm/proxy/policy_engine/attachment_registry.py @@ -30,6 +30,23 @@ class PolicyAttachmentMatch(TypedDict): matched_via: str +def _attachment_specificity(attachment: PolicyAttachment) -> tuple[int, int]: + if attachment.is_global(): + return (0, 0) + + dims: Final = tuple( + specificity + for values, specificity in ( + (attachment.teams, 1), + (attachment.keys, 2), + (attachment.tags, 3), + (attachment.models, 4), + ) + if values + ) + return (max(dims, default=0), len(dims)) + + class AttachmentRegistry: """ In-memory registry for storing and managing policy attachments. @@ -116,31 +133,26 @@ class AttachmentRegistry: """ from litellm.proxy.policy_engine.policy_matcher import PolicyMatcher - results: Final[list[PolicyAttachmentMatch]] = [] - seen_policies: Final[set[str]] = set() + matching_attachments: Final = sorted( + ( + attachment + for attachment in self._attachments + if PolicyMatcher.scope_matches(scope=attachment.to_policy_scope(), context=context) + ), + key=_attachment_specificity, + ) + unique_attachments: Final = tuple( + next(attachment for attachment in matching_attachments if attachment.policy == policy_name) + for policy_name in dict.fromkeys(attachment.policy for attachment in matching_attachments) + ) - for attachment in self._attachments: - scope = attachment.to_policy_scope() - if PolicyMatcher.scope_matches(scope=scope, context=context): - if attachment.policy not in seen_policies: - seen_policies.add(attachment.policy) - matched_via = self._describe_match_reason(attachment, context) - results.append( - { - "policy_name": attachment.policy, - "matched_via": matched_via, - } - ) - verbose_proxy_logger.debug( - "Attachment matched: policy=%s, matched_via=%s, context=(team=%s, key=%s, model=%s)", - attachment.policy, - matched_via, - context.team_alias, - context.key_alias, - context.model, - ) - - return results + return [ + { + "policy_name": attachment.policy, + "matched_via": self._describe_match_reason(attachment, context), + } + for attachment in unique_attachments + ] @staticmethod def _describe_match_reason(attachment: PolicyAttachment, context: PolicyMatchContext) -> str: diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 94e74b20297..c8fd7edfed6 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -304,7 +304,7 @@ from litellm.litellm_core_utils.sensitive_data_masker import ( ) from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.vertex_ai.vertex_llm_base import VertexBase -from litellm.proxy._lazy_features import attach_lazy_features +from litellm.proxy._lazy_features import attach_lazy_features, reserve_lazy_slot from litellm.proxy._types import * from litellm.proxy.analytics_endpoints.analytics_endpoints import ( router as analytics_router, @@ -639,13 +639,8 @@ from litellm.proxy.openai_files_endpoints.files_endpoints import ( from litellm.proxy.openai_files_endpoints.files_endpoints import ( set_files_config, ) -from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - openai_passthrough_router, - passthrough_endpoint_router, - vertex_ai_live_websocket_passthrough, -) -from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( - router as llm_passthrough_router, +from litellm.proxy.pass_through_endpoints.openai_passthrough_endpoints import ( + router as openai_passthrough_router, ) from litellm.proxy.pass_through_endpoints.pass_through_endpoints import ( initialize_pass_through_endpoints, @@ -660,6 +655,7 @@ from litellm.proxy.rag_endpoints.endpoints import router as rag_router from litellm.proxy.rerank_endpoints.endpoints import router as rerank_router from litellm.proxy.response_api_endpoints.endpoints import router as response_router from litellm.proxy.route_llm_request import route_request +from litellm.proxy.route_priority import hot_routes_first from litellm.proxy.search_endpoints.endpoints import router as search_router from litellm.proxy.shutdown.graceful_shutdown_manager import GracefulShutdownManager from litellm.proxy.spend_tracking.budget_reservation import get_budget_window_start @@ -5813,6 +5809,16 @@ class ProxyConfig: default_redis_ttl=ttl, ) + ### USER API KEY CACHE MAX SIZE (in-memory tier shared by keys, teams, users, end users, ...) ### + if "user_api_key_cache_max_size" in general_settings: + user_api_key_cache.update_in_memory_max_size( + ConfigGeneralSettings.model_validate( + MappingProxyType( + {"user_api_key_cache_max_size": general_settings["user_api_key_cache_max_size"]} + ) + ).user_api_key_cache_max_size + ) + ### PKCE MULTI-INSTANCE PREREQUISITE CHECK ### # PKCE verifiers are stored in redis_usage_cache when available so they can # be read back by any instance (not just the one that started the auth flow). @@ -6047,6 +6053,10 @@ class ProxyConfig: set_files_config(config=files_config) ## default config for vertex ai routes + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + passthrough_endpoint_router, + ) + default_vertex_config: Final = config.get("default_vertex_config", None) passthrough_endpoint_router.set_default_vertex_config(config=default_vertex_config) @@ -7059,6 +7069,23 @@ class ProxyConfig: "enable_openai_websocket_passthrough" ) + if "user_api_key_cache_max_size" not in self._yaml_general_settings_keys: + db_cache_max_size: Final = _general_settings.get("user_api_key_cache_max_size") + try: + cache_max_size: Final = ConfigGeneralSettings.model_validate( + MappingProxyType({"user_api_key_cache_max_size": db_cache_max_size}) + ).user_api_key_cache_max_size + except ValidationError: + verbose_proxy_logger.warning( + "Ignoring invalid general_settings.user_api_key_cache_max_size=%r from the DB", db_cache_max_size + ) + else: + if cache_max_size is None: + general_settings.pop("user_api_key_cache_max_size", None) + else: + general_settings["user_api_key_cache_max_size"] = cache_max_size + user_api_key_cache.update_in_memory_max_size(cache_max_size) + ## STORE MODEL IN DB ## if "store_model_in_db" in _general_settings: value = _general_settings["store_model_in_db"] @@ -10830,14 +10857,14 @@ async def model_info( # Use the actual litellm model from the deployment to get provider info _, provider, _, _ = litellm.get_llm_provider(model=deployment.litellm_params.model) - response_id: Final = internal_to_public.get(resolved_model_id, model_id) - return create_model_info_response( - model_id=response_id, + response: Final = create_model_info_response( + model_id=resolved_model_id, provider=provider, include_metadata=False, fallback_type=None, llm_router=llm_router, ) + return {**response, "id": internal_to_public.get(resolved_model_id, model_id)} # mutable-ok: response id differs def _blocked_response_usage(original_response: object | None) -> "litellm.Usage": @@ -11763,6 +11790,10 @@ async def vertex_ai_live_passthrough_endpoint( This endpoint delegates to the WebSocket function defined in llm_passthrough_endpoints.py """ + from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( + vertex_ai_live_websocket_passthrough, + ) + return await vertex_ai_live_websocket_passthrough( websocket=websocket, model=model, @@ -16967,6 +16998,7 @@ _GENERAL_SETTINGS_CONFIG_LIST_FIELD_TYPES: Final[Mapping[str, str]] = MappingPro "cancel_on_disconnect": "Boolean", "disable_auto_add_proxy_admin_to_teams": "Boolean", "apply_user_budget_to_team_keys": "Boolean", + "user_api_key_cache_max_size": "Integer", } ) @@ -18668,7 +18700,7 @@ app.include_router(credential_router) app.include_router(openai_passthrough_router) app.include_router(batches_router) app.include_router(openai_files_router) -app.include_router(llm_passthrough_router) +reserve_lazy_slot(app, "llm_passthrough") app.include_router(pass_through_router) app.include_router(health_router) app.include_router(key_management_router) @@ -18708,6 +18740,7 @@ app.include_router(ui_discovery_endpoints_router) app.include_router(google_router) attach_lazy_features(app) +app.router.routes = hot_routes_first(app.router.routes) app.add_middleware( RequestSizeLimitMiddleware, get_max_request_size_mb=lambda: general_settings.get("max_request_size_mb"), diff --git a/litellm/proxy/route_priority.py b/litellm/proxy/route_priority.py new file mode 100644 index 00000000000..77815b678a6 --- /dev/null +++ b/litellm/proxy/route_priority.py @@ -0,0 +1,24 @@ +"""Starlette matches routes in registration order, so the routes that take the most traffic go first.""" + +from collections.abc import Sequence +from typing import Final + +from starlette.routing import BaseRoute, Route + +HOT_ROUTE_PATHS: Final[frozenset[str]] = frozenset( + ( + "/health/liveliness", + "/health/liveness", + "/v1/chat/completions", + "/chat/completions", + "/v1/messages", + ) +) + + +def _is_hot(route: BaseRoute) -> bool: + return isinstance(route, Route) and route.path in HOT_ROUTE_PATHS + + +def hot_routes_first(routes: Sequence[BaseRoute]) -> list[BaseRoute]: # mutable-ok: assigned to Router.routes, a list + return sorted(routes, key=lambda route: not _is_hot(route)) diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 817df082d8c..7d521d54791 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1514,6 +1514,7 @@ model LiteLLM_AutoRouterSession { classifier_cost Float @default(0) classifier_cost_recorded_turns Int @default(0) tier_turns Json @default("{}") + baseline_models Json @default("{}") @@id([api_key, session_id, router_name]) @@index([last_turn_at], map: "idx_autorouter_session_last_turn") diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index da79328fa59..8cfb6354dd0 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -2925,6 +2925,32 @@ async def _fetch_session_representatives( return [rep_by_key[key] for key in session_keys if key in rep_by_key] # mutable-ok: rows are enriched in place +async def _count_grouped_sessions( + prisma_client: "PrismaClient", + where_clause: str, + sql_params: Sequence[object], + next_param_index: int, +) -> tuple[int, bool]: + """Count the sessions matching the filter, returning ``(total, total_is_capped)`` bounded by the count cap.""" + count_query: Final = f""" + SELECT COUNT(*) AS total_count + FROM ( + SELECT 1 + FROM "LiteLLM_SpendLogs" + WHERE {where_clause} + GROUP BY {_SESSION_GROUP_KEY_SQL} + LIMIT ${next_param_index} + ) AS bounded_sessions + """ + count_rows: Final[Sequence[_SpendLogsCountRow]] = await _query_raw( + prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1 + ) + raw_total: Final = int(count_rows[0]["total_count"]) if count_rows else 0 + return ( + (SPEND_LOGS_PAGINATION_COUNT_CAP, True) if raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP else (raw_total, False) + ) + + async def _ui_session_grouped_spend_logs( prisma_client: "PrismaClient", sql_conditions: Sequence[str], @@ -2944,11 +2970,19 @@ async def _ui_session_grouped_spend_logs( next ``page_size`` sessions ordered by ``(MAX(startTime), session_key, api_key)``, resumed from the ``session_cursor`` keyset ``'||'`` instead of an OFFSET, so - page depth does not degrade the query plan. Each session is represented + page depth does not degrade the query plan. A request for ``page > 1`` + without a cursor (the UI jumping straight to the last page, or back to a + page it never walked through) falls back to ``OFFSET (page - 1) * + page_size``, trimmed to the end of the ``SPEND_LOGS_PAGINATION_COUNT_CAP`` + window the capped ``total`` promises, so a page never runs past that total + and one starting at or past it returns no rows without a query. Each session is represented by its newest non-MCP row, enriched by ``_build_ui_spend_logs_response`` exactly like the flat listing, and the response carries ``next_session_cursor`` / ``has_more`` while ``total`` counts sessions - (capped like the flat total). + (capped like the flat total). A page that runs out of sessions while still + holding some is itself the end of the list, so its ``total`` is + ``offset + len(page)`` and the grouped count query is skipped; a page that + starts past the end says nothing about the total, so that one is counted. """ where_clause: Final = " AND ".join(sql_conditions) if sql_conditions else "TRUE" cmp_op: Final = "<" if sort_desc else ">" @@ -2963,6 +2997,10 @@ async def _ui_session_grouped_spend_logs( ) cursor_params: Final[tuple[object, ...]] = cursor if cursor else () limit_index: Final = next_param_index + len(cursor_params) + offset: Final = (page - 1) * page_size if cursor is None else 0 + page_limit: Final = min(page_size, SPEND_LOGS_PAGINATION_COUNT_CAP - offset) + offset_params: Final[tuple[int, ...]] = (offset,) if offset and page_limit > 0 else () + offset_clause: Final = f"OFFSET ${limit_index + 1}" if offset_params else "" page_query: Final = f""" SELECT {_SESSION_KEY_EXPR} AS session_key, @@ -2973,36 +3011,29 @@ async def _ui_session_grouped_spend_logs( GROUP BY {_SESSION_GROUP_KEY_SQL} {having_clause} ORDER BY MAX("startTime") {direction}, {_SESSION_KEY_EXPR} {direction}, api_key {direction} - LIMIT ${limit_index} + LIMIT ${limit_index} {offset_clause} """ - page_rows: Final[Sequence[_SessionPageRow]] = await _query_raw( - prisma_client, page_query, *sql_params, *cursor_params, page_size + 1 + page_rows: Final[Sequence[_SessionPageRow]] = ( + () + if page_limit <= 0 + else await _query_raw(prisma_client, page_query, *sql_params, *cursor_params, page_limit + 1, *offset_params) ) - has_more: Final = len(page_rows) > page_size - visible_rows: Final = page_rows[:page_size] + has_more: Final = len(page_rows) > page_limit + visible_rows: Final = page_rows[:page_limit] next_cursor: Final = ( f"{visible_rows[-1]['last_activity']}|{visible_rows[-1]['api_key']}|{visible_rows[-1]['session_key']}" if has_more and visible_rows else None ) - count_query: Final = f""" - SELECT COUNT(*) AS total_count - FROM ( - SELECT 1 - FROM "LiteLLM_SpendLogs" - WHERE {where_clause} - GROUP BY {_SESSION_GROUP_KEY_SQL} - LIMIT ${next_param_index} - ) AS bounded_sessions - """ - count_rows: Final[Sequence[_SpendLogsCountRow]] = await _query_raw( - prisma_client, count_query, *sql_params, SPEND_LOGS_PAGINATION_COUNT_CAP + 1 + page_starts_inside_the_list: Final = offset == 0 or len(page_rows) > 0 + page_ends_the_list: Final = cursor is None and page_limit > 0 and not has_more and page_starts_inside_the_list + total_records, total_is_capped = ( + (offset + len(page_rows), False) + if page_ends_the_list + else await _count_grouped_sessions(prisma_client, where_clause, sql_params, next_param_index) ) - raw_total: Final = int(count_rows[0]["total_count"]) if count_rows else 0 - total_is_capped: Final = raw_total > SPEND_LOGS_PAGINATION_COUNT_CAP - total_records: Final = SPEND_LOGS_PAGINATION_COUNT_CAP if total_is_capped else raw_total session_keys: Final = tuple((row["session_key"], row["api_key"]) for row in visible_rows) data: Final[list[dict[str, object]]] = ( # mutable-ok: _build_ui_spend_logs_response writes onto each row diff --git a/litellm/proxy/spend_tracking/spend_tracking_utils.py b/litellm/proxy/spend_tracking/spend_tracking_utils.py index f0f38358cf0..01398d38687 100644 --- a/litellm/proxy/spend_tracking/spend_tracking_utils.py +++ b/litellm/proxy/spend_tracking/spend_tracking_utils.py @@ -19,6 +19,7 @@ from litellm.constants import ( LITTELM_CLI_SERVICE_ACCOUNT_NAME, LITTELM_INTERNAL_HEALTH_SERVICE_ACCOUNT_NAME, MAX_SPEND_LOG_MODEL_NAME_LENGTH, + MCP_SPEND_LOG_MODEL_PREFIX, REDACTED_BY_LITELM_STRING, SESSION_ID_OMITTED_METADATA_KEY, UNKNOWN_MODEL_SPEND_LOG_MODEL, @@ -338,7 +339,8 @@ def _sl_attribution_fallback( def _looks_like_model_name(model: str) -> bool: - return len(model) <= MAX_SPEND_LOG_MODEL_NAME_LENGTH and not any(char.isspace() for char in model) + candidate: Final = model.removeprefix(MCP_SPEND_LOG_MODEL_PREFIX) + return len(candidate) <= MAX_SPEND_LOG_MODEL_NAME_LENGTH and not any(char.isspace() for char in candidate) def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogsPayload: diff --git a/litellm/proxy/utils.py b/litellm/proxy/utils.py index 253494b02f4..8e6142090a4 100644 --- a/litellm/proxy/utils.py +++ b/litellm/proxy/utils.py @@ -877,9 +877,10 @@ _EMPTY_LIFT: Final = MappingProxyType({}) def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, object]: """Failure-path callbacks run after ``litellm_logging_obj`` is popped from request_data (it is not serialisable), so the caller merges these fields - onto request_data first: the first-handoff instant for preprocessing - latency, recovered or estimated usage for token counts, and the standard - logging object for deployment attribution on failed-request spend logs.""" + onto request_data first: the request start and first-handoff instants for + duration and preprocessing latency, the call type, recovered or estimated + usage for token counts, and the standard logging object for deployment + attribution on failed-request spend logs.""" _logging_obj: Final = request_data.get("litellm_logging_obj") if _logging_obj is None: return _EMPTY_LIFT @@ -891,7 +892,9 @@ def _failure_fields_to_lift(request_data: Mapping[str, object]) -> Mapping[str, dispatched=_first_handoff is not None, ) _entries: Final = ( + ("start_time", _model_call_details.get("start_time")), ("first_api_call_start_time", _first_handoff), + ("call_type", _model_call_details.get("call_type")), ("combined_usage_object", None if _usage_to_lift is None else _usage_to_lift[0]), ("response_cost", None if _usage_to_lift is None else (_usage_to_lift[1] or 0.0)), ("standard_logging_object", _model_call_details.get("standard_logging_object")), diff --git a/litellm/repositories/__init__.py b/litellm/repositories/__init__.py index 881f7a66cea..dcf9ddfc32a 100644 --- a/litellm/repositories/__init__.py +++ b/litellm/repositories/__init__.py @@ -2,6 +2,7 @@ Repository classes for database operations. """ +from litellm.repositories.autorouter_session_repository import AutoRouterSessionRepository from litellm.repositories.budget_repository import BudgetRepository from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.credentials_repository import CredentialsRepository @@ -92,6 +93,7 @@ __all__ = [ "AdaptiveRouterStateRepository", "AgentsRepository", "AuditLogRepository", + "AutoRouterSessionRepository", "BatchTable", "BudgetCascadeUnitOfWork", "BudgetRepository", diff --git a/litellm/repositories/autorouter_session_repository.py b/litellm/repositories/autorouter_session_repository.py new file mode 100644 index 00000000000..d05ef9421ca --- /dev/null +++ b/litellm/repositories/autorouter_session_repository.py @@ -0,0 +1,34 @@ +""" +Repository for the auto-router per-session rollup (LiteLLM_AutoRouterSession). +""" + +from typing import TYPE_CHECKING, Final + +from litellm.models.autorouter_session import LiteLLM_AutoRouterSession +from litellm.repositories.base_repository import BaseRepository +from litellm.repositories.prisma_protocols import TableActions + +if TYPE_CHECKING: + from prisma import models as prisma_models + + +class AutoRouterSessionRepository(BaseRepository[LiteLLM_AutoRouterSession]): + @property + def table(self) -> TableActions["prisma_models.LiteLLM_AutoRouterSession"]: + return self.prisma_client.db.litellm_autoroutersession + + @property + def model_class(self) -> type[LiteLLM_AutoRouterSession]: + return LiteLLM_AutoRouterSession + + async def find_latest_for_key(self, api_key: str, session_id: str) -> LiteLLM_AutoRouterSession | None: + """The session's most recently active router row under exactly this key hash, or None. + + The key is the row's own partition, not a filter over a wider read: the spend writer keyed the + row under the caller's api_key, so a key can only ever see what it wrote itself. + """ + record: Final = await self.table.find_first( + where={"api_key": api_key, "session_id": session_id}, # mutable-ok: Prisma where filter must be a dict + order={"last_turn_at": "desc"}, # mutable-ok: Prisma order clause must be a dict + ) + return self._to_model(record) diff --git a/litellm/router.py b/litellm/router.py index 9e6db66db83..597bcfaa20f 100644 --- a/litellm/router.py +++ b/litellm/router.py @@ -12669,6 +12669,7 @@ class Router: input: str | list | None = None, specific_deployment: bool | None = False, parent_otel_span: Span | None = None, + health_check_probe: bool = False, ) -> list[dict] | dict: """ Get the healthy deployments for a model. @@ -12718,6 +12719,7 @@ class Router: healthy_deployments = await self._async_filter_health_check_unhealthy_deployments( healthy_deployments=healthy_deployments, parent_otel_span=parent_otel_span, + health_check_probe=health_check_probe, ) cooldown_deployments: Final = await _async_get_cooldown_deployments( @@ -14100,6 +14102,7 @@ class Router: self, healthy_deployments: list[dict], parent_otel_span: Span | None = None, + health_check_probe: bool = False, ) -> list[dict]: """ Filter out deployments marked unhealthy by background health checks. @@ -14136,8 +14139,7 @@ class Router: ] if not filtered: - verbose_router_logger.warning("All deployments marked unhealthy by health checks, bypassing health filter") - return healthy_deployments + return [] if health_check_probe else healthy_deployments # mutable-ok: empty list signals unavailable probe return filtered diff --git a/litellm/router_strategy/adaptive_router/adaptive_router.py b/litellm/router_strategy/adaptive_router/adaptive_router.py index 12ccacbbc1d..7f376b46a8d 100644 --- a/litellm/router_strategy/adaptive_router/adaptive_router.py +++ b/litellm/router_strategy/adaptive_router/adaptive_router.py @@ -461,6 +461,8 @@ class AdaptiveRouter: if d_alpha == 0 and d_beta == 0: continue cell_key = (attribution_type, target_model) + if cell_key not in self._cells: + continue self._cells[cell_key] = apply_delta( self._cells[cell_key], d_alpha, diff --git a/litellm/router_strategy/complexity_router/README.md b/litellm/router_strategy/complexity_router/README.md index 1a4764c291e..4aa342ea59f 100644 --- a/litellm/router_strategy/complexity_router/README.md +++ b/litellm/router_strategy/complexity_router/README.md @@ -270,6 +270,18 @@ change or default takeover records `cause: modality_escalation` with the displac pinned by session affinity, and by default a KEPT session pin bypasses the gate: a session pinned to a text-only model keeps it even when an image arrives. +Context-window and modality recovery take priority over the default model. If a compatible tier +cannot serve, the router checks the remaining compatible recovery tiers before using `default_model`. +A capacity failure without those constraints tries the selected tier's peers, then the default + +The default must fit the context and accept the request's modality. It cannot bypass routing plugins +or a plan-mode floor. Context fit uses the auto-router's existing buffer even when Router-wide pre-call +checks are off. Missing context metadata retains the existing unknown-window behavior + +Health fallback records `cause: health_default_fallback` and `health_displaced:` in `signals`. +It does not replace the session's tier pin. Adaptive feedback retains the model that actually served, +but a default outside the adaptive candidate pool does not become a normal candidate + Add `modality_pin_override: true` to lift that last exemption. The image turn is then re-placed the same way every other decision is, and records `cause: modality_pin_override` whether or not the tier moved, since the model left the pin either way. The pin itself is untouched: the session diff --git a/litellm/router_strategy/complexity_router/complexity_router.py b/litellm/router_strategy/complexity_router/complexity_router.py index 7519f4c5156..9deccc9a468 100644 --- a/litellm/router_strategy/complexity_router/complexity_router.py +++ b/litellm/router_strategy/complexity_router/complexity_router.py @@ -937,6 +937,7 @@ def _decision_is_pinnable(decision: StandardLoggingRoutingDecision | None) -> bo "modality_escalation", "modality_pin_override", "health_failover", + "health_default_fallback", ) and not decision.get("context_escalated") and _CLASSIFIER_CIRCUIT_OPEN_SIGNAL not in (decision.get("signals") or ()) @@ -1098,6 +1099,15 @@ def _group_provably_fits(facts: tuple[int | None, bool], needed: int, buffer: fl return window is not None and not has_unknown and needed <= int(window * buffer) +class _RequestContextFit(NamedTuple): + facts: Mapping[str, tuple[int | None, bool]] + needed: int | None + buffer: float + + def accepts(self, model: str) -> bool: + return self.needed is None or _window_can_hold(self.facts.get(model, (None, True))[0], self.needed, self.buffer) + + class _ContextWindowPlacement(NamedTuple): """Where the context-window gate placed the request: the placement tier, the subset of its pool the pick may use, and every configured group not provably misfit (the adaptive filter).""" @@ -2681,12 +2691,32 @@ class ComplexityRouter(CustomLogger): verbose_router_logger.debug("ComplexityRouter: context-window token count failed. Got - %s", e) return None + async def _request_context_fit( + self, + resolved_messages: Sequence[Mapping[str, object]] | None, + request_kwargs: Mapping[str, object], + ) -> _RequestContextFit: + if not self.config.enable_context_window_escalation or not resolved_messages: + return _RequestContextFit(EMPTY_MAPPING, None, self.config.context_window_escalation_buffer) + names: Final = frozenset(model for pool in self._tier_pools().values() for model in pool) | frozenset( + (self.config.default_model,) if self.config.default_model else () + ) + facts: Final = MappingProxyType({name: self._group_window_facts(name) for name in names}) + known: Final = tuple(window for window, _ in facts.values() if window is not None) + buffer: Final = self.config.context_window_escalation_buffer + needs_count: Final = known and self._request_byte_upper_bound(resolved_messages, request_kwargs) > int( + min(known) * buffer + ) + needed: Final = await self._counted_request_tokens(resolved_messages, request_kwargs) if needs_count else None + return _RequestContextFit(facts=facts, needed=needed, buffer=buffer) + async def _context_window_placement( self, tier: ComplexityTier | str, resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: Mapping[str, object], pool_override: tuple[str, ...] | None = None, + context_fit: _RequestContextFit | None = None, ) -> _ContextWindowPlacement | None: """Correct a decided placement whose models provably cannot hold the prompt, or None (the placement stands). Only a real tokenizer count ever moves a request, escalation @@ -2698,17 +2728,10 @@ class ComplexityRouter(CustomLogger): pool: Final = pool_override if pool_override is not None else tuple(pools.get(_tier_name(tier), ())) if not pool: return None - facts: Final = MappingProxyType({group: self._group_window_facts(group) for group in pool}) - known_windows: Final = tuple(window for window, _ in facts.values() if window is not None) - if not known_windows: + fit: Final = context_fit or await self._request_context_fit(resolved_messages, request_kwargs) + if fit.needed is None: return None - buffer: Final = self.config.context_window_escalation_buffer - if self._request_byte_upper_bound(resolved_messages, request_kwargs) <= int(min(known_windows) * buffer): - return None - needed: Final = await self._counted_request_tokens(resolved_messages, request_kwargs) - if needed is None: - return None - return self._placement_for_tokens(tier=tier, pool=pool, pools=pools, facts=facts, needed=needed) + return self._placement_for_tokens(tier=tier, pool=pool, pools=pools, facts=fit.facts, needed=fit.needed) def _placement_for_tokens( self, @@ -2720,14 +2743,16 @@ class ComplexityRouter(CustomLogger): needed: int, ) -> _ContextWindowPlacement | None: buffer: Final = self.config.context_window_escalation_buffer - in_tier: Final = tuple(group for group in pool if _window_can_hold(facts[group][0], needed, buffer)) + in_tier: Final = tuple( + group for group in pool if _window_can_hold(facts.get(group, (None, True))[0], needed, buffer) + ) if in_tier and len(in_tier) == len(pool): return None holdable: Final = frozenset( group for tier_pool in pools.values() for group in tier_pool - if _window_can_hold(self._group_window_facts(group)[0], needed, buffer) + if _window_can_hold(facts.get(group, (None, True))[0], needed, buffer) ) if in_tier: return _ContextWindowPlacement(tier=tier, allowed_models=in_tier, holdable_models=holdable) @@ -2735,7 +2760,7 @@ class ComplexityRouter(CustomLogger): proven = tuple( group for group in pools.get(name, ()) - if _group_provably_fits(self._group_window_facts(group), needed, buffer) + if _group_provably_fits(facts.get(group, (None, True)), needed, buffer) ) if proven: return _ContextWindowPlacement( @@ -2881,6 +2906,7 @@ class ComplexityRouter(CustomLogger): messages: list[dict[str, Any]] | None, # mutable-ok: forwarded verbatim to the list-typed re-pick resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: dict, # mutable-ok: same shape the hook receives + context_fit: _RequestContextFit | None = None, ) -> PreRoutingHookResponse: """Replace a routed model that cannot accept this request's image input. @@ -2911,7 +2937,8 @@ class ComplexityRouter(CustomLogger): or self._model_accepts_image_input(response.model) ): return response - eligible: Final = self._modality_eligible_models() + fit: Final = context_fit or await self._request_context_fit(resolved_messages, request_kwargs) + eligible: Final = frozenset(name for name in self._modality_eligible_models() if fit.accepts(name)) names: Final = self.config.tier_names() pools: Final = self._tier_pools() decided: Final = decision.get("tier") if decision is not None else None @@ -3030,12 +3057,15 @@ class ComplexityRouter(CustomLogger): Every way the owner says "nothing here can serve this" is a negative verdict: no healthy deployment for the group at all (BadRequestError, which ContextWindowExceededError - subclasses), every deployment filtered out (RouterRateLimitError), and every deployment - over its RPM (RouterRateLimitErrorBasic). Anything else is unknown rather than negative, - so it reads as capacity: absent information must never decide the verdict. + subclasses), every deployment filtered out (RouterRateLimitError), every deployment over + its RPM (RouterRateLimitErrorBasic), and every deployment refused by a filter that reports + exhaustion as a bare ValueError naming a RouterErrors marker -- provider and deployment + budgets, and tag routing, which have no typed error of their own. Anything else is unknown + rather than negative, so it reads as capacity: absent information must never decide the + verdict. """ from litellm.exceptions import BadRequestError - from litellm.types.router import RouterRateLimitError, RouterRateLimitErrorBasic + from litellm.types.router import RouterErrors, RouterRateLimitError, RouterRateLimitErrorBasic probe_kwargs: Final = dict(request_kwargs) # mutable-ok: the owner pops routing keys off the dict it is handed try: @@ -3045,10 +3075,15 @@ class ComplexityRouter(CustomLogger): messages=messages, input=input, parent_otel_span=_get_parent_otel_span_from_kwargs(request_kwargs), + health_check_probe=True, ) - except (RouterRateLimitError, RouterRateLimitErrorBasic, BadRequestError): + except (RouterRateLimitError, RouterRateLimitErrorBasic, BadRequestError) as exc: + verbose_router_logger.debug("health probe unavailable model=%s error=%s", model_name, type(exc).__name__) return False except Exception as exc: # noqa: BLE001 # a speculative eligibility read must fail open on unknown faults + if isinstance(exc, ValueError) and any(marker.value in str(exc) for marker in RouterErrors): + verbose_router_logger.debug("health probe exhausted model=%s error=%s", model_name, exc) + return False verbose_router_logger.debug( "ComplexityRouter: eligibility probe for %s failed, treating the group as live: %s", model_name, exc ) @@ -3062,76 +3097,124 @@ class ComplexityRouter(CustomLogger): input: str | list | None, # mutable-ok: mirrors the owner's own input parameter, which this forwards verbatim resolved_messages: Sequence[Mapping[str, object]] | None, request_kwargs: dict, # mutable-ok: same shape the hook receives + context_fit: _RequestContextFit | None = None, ) -> PreRoutingHookResponse: - """Replace a decided model group that has no serving capacity with a live peer in the same tier. - - Applied to the decided response at the hook's exits, so every arm that can place a request - is covered by one owner: a fresh classification, a replayed or escalated session pin, a - plan-mode floor, a context-window escalation, an adaptive pick, and whatever arm is added - next. Peers come from the DECIDED tier only; climbing to another tier is deliberately not - done here, since a higher tier costs more than the classifier asked for. - - Serving capacity is one question asked of one owner (`_model_group_can_serve`), so the - substitute is only ever a group the pipeline would actually accept for this request. The - pick then runs through `_pick_model_for_tier`, so routing plugins decide the substitute - exactly as they decided the original. - - Fails open everywhere it cannot be sure: an unreadable eligibility view, a decision - carrying no tier (default_model), or a tier whose every peer is unusable too. It fails - CLOSED on a plugin that empties the pool, leaving the original decision to fail rather - than serving a model the plugin excluded. - """ + """Try compatible tier recovery before the default, preserving request policy and fit.""" decision: Final = response.routing_decision decided_tier: Final = decision.get("tier") if decision is not None else None if decision is None or not isinstance(decided_tier, str): return response - peers: Final = tuple(self._tier_pools().get(decided_tier, ())) - if len(peers) < 2: - return response - if await self._model_group_can_serve(response.model, messages, input, request_kwargs): + fit: Final = context_fit or await self._request_context_fit(resolved_messages, request_kwargs) + if fit.accepts(response.model) and await self._model_group_can_serve( + response.model, messages, input, request_kwargs + ): return response eligible: Final = ( self._modality_eligible_models() if self.config.modality_routing and resolved_messages and request_contains_image_content(resolved_messages) else None ) - candidates: Final = tuple( - peer for peer in peers if peer != response.model and (eligible is None or peer in eligible) + pools: Final = self._tier_pools() + context_recovery: Final = bool(decision.get("context_escalated")) or any( + not fit.accepts(model) for model in pools.get(decided_tier, ()) ) - if not candidates: - return response - servable: Final = await asyncio.gather( - *(self._model_group_can_serve(peer, messages, input, request_kwargs) for peer in candidates) + modality_recovery: Final = eligible is not None + names: Final = self.config.tier_names() + tiers: Final = ( + tuple(names[names.index(decided_tier) :]) + if (context_recovery or modality_recovery) and decided_tier in names + else (decided_tier,) ) - live: Final = tuple(peer for peer, can_serve in zip(candidates, servable) if can_serve) - if not live: - return response - repick_messages: Final = ( - list(resolved_messages) if resolved_messages else None # mutable-ok: the pick's param is list-typed - ) - try: - new_model: Final = await self._pick_model_for_tier( - decided_tier if self.config.has_custom_tiers else ComplexityTier(decided_tier), - messages, - repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them - request_kwargs, - allowed_models=live, + + async def recover_tier(candidate_tier: str) -> PreRoutingHookResponse | None: + peers: Final = tuple( + model + for model in pools.get(candidate_tier, ()) + if not context_recovery + or candidate_tier == decided_tier + or fit.needed is None + or _group_provably_fits(fit.facts.get(model, (None, True)), fit.needed, fit.buffer) ) - except ValueError as exc: - verbose_router_logger.debug( - "ComplexityRouter: health failover found no candidate the routing plugins allow: %s", exc + candidates: Final = tuple( + peer + for peer in peers + if peer != response.model and fit.accepts(peer) and (eligible is None or peer in eligible) ) + servable: Final = await asyncio.gather( + *(self._model_group_can_serve(peer, messages, input, request_kwargs) for peer in candidates) + ) + live: Final = tuple(peer for peer, can_serve in zip(candidates, servable) if can_serve) + if live: + repick_messages: Final = ( + list(resolved_messages) if resolved_messages else None # mutable-ok: the pick's param is list-typed + ) + try: + new_model: Final = await self._pick_model_for_tier( + candidate_tier if self.config.has_custom_tiers else ComplexityTier(candidate_tier), + messages, + repick_messages, # pyright: ignore[reportArgumentType] # hook-resolved message dicts; the pick only reads them + request_kwargs, + allowed_models=live, + ) + except ValueError as exc: + verbose_router_logger.debug( + "ComplexityRouter: health failover found no candidate the routing plugins allow: %s", exc + ) + else: + self._restamp_adaptive_choice(request_kwargs, response.model, new_model) + verbose_router_logger.info( + "ComplexityRouter: routing decision cause=health_failover, routed_model=%s, displaced=%s", + new_model, + response.model, + ) + new_decision: Final = self._build_routing_decision( + routed_model=new_model, + cause="health_failover", + tier=candidate_tier, + score=decision.get("score"), + signals=(*(decision.get("signals") or ()), f"health_displaced:{response.model}"), + matched_keyword=decision.get("matched_keyword"), + escalation_keyword=decision.get("escalation_keyword"), + escalated=bool(decision.get("escalated", False)), + classifier_model=decision.get("classifier_model"), + classifier_cost=decision.get("classifier_cost"), + conversation_continuing=bool(decision.get("conversation_continuing", True)), + tier_litellm_params=self._litellm_params_for_model(candidate_tier, new_model), + context_escalation_original_tier=decision.get("context_escalation_original_tier"), + ) + return response.model_copy( + update={ # mutable-ok: model_copy types update as a plain dict + "model": new_model, + "litellm_params": self._litellm_params_for_model(candidate_tier, new_model), + "routing_decision": new_decision, + } + ) + return None + + for candidate_tier in tiers: + if (recovered := await recover_tier(candidate_tier)) is not None: + return recovered + default_model: Final = self.config.default_model + plan_mode_active: Final = self._matched_plan_mode_signal(request_kwargs, resolved_messages) is not None + if ( + plan_mode_active + or self.config.plugins + or not default_model + or default_model == response.model + or not fit.accepts(default_model) + or (eligible is not None and default_model not in eligible) + or not await self._model_group_can_serve(default_model, messages, input, request_kwargs) + ): return response - self._restamp_adaptive_choice(request_kwargs, response.model, new_model) + self._restamp_adaptive_choice(request_kwargs, response.model, default_model) verbose_router_logger.info( - "ComplexityRouter: routing decision cause=health_failover, routed_model=%s, displaced=%s", - new_model, + "ComplexityRouter: routing decision cause=health_default_fallback, routed_model=%s, displaced=%s", + default_model, response.model, ) - new_decision: Final = self._build_routing_decision( - routed_model=new_model, - cause="health_failover", - tier=decision.get("tier"), + default_decision: Final = self._build_routing_decision( + routed_model=default_model, + cause="health_default_fallback", score=decision.get("score"), signals=(*(decision.get("signals") or ()), f"health_displaced:{response.model}"), matched_keyword=decision.get("matched_keyword"), @@ -3140,14 +3223,14 @@ class ComplexityRouter(CustomLogger): classifier_model=decision.get("classifier_model"), classifier_cost=decision.get("classifier_cost"), conversation_continuing=bool(decision.get("conversation_continuing", True)), - tier_litellm_params=self._litellm_params_for_model(decided_tier, new_model), + tier_litellm_params=self._litellm_params_for_model(None, default_model), context_escalation_original_tier=decision.get("context_escalation_original_tier"), ) return response.model_copy( update={ # mutable-ok: model_copy types update as a plain dict - "model": new_model, - "litellm_params": self._litellm_params_for_model(decided_tier, new_model), - "routing_decision": new_decision, + "model": default_model, + "litellm_params": self._litellm_params_for_model(None, default_model), + "routing_decision": default_decision, } ) @@ -3462,6 +3545,7 @@ class ComplexityRouter(CustomLogger): # chat-completions messages, so it is real work on every non-chat surface, and # both the conversation shape and the classifier read the same list. resolved_messages: Final = self._resolve_messages(messages, request_kwargs) + context_fit: Final = await self._request_context_fit(resolved_messages, request_kwargs) marker_pairs: Final = self._reminder_markers_for_request(request_kwargs) conversation_continuing: Final = _conversation_is_continuing(resolved_messages) @@ -3512,7 +3596,11 @@ class ComplexityRouter(CustomLogger): pin_source_tier: Final = self._tier_for_model(routed_model) pin_placement: Final = ( await self._context_window_placement( - pin_source_tier, resolved_messages, request_kwargs, pool_override=(routed_model,) + pin_source_tier, + resolved_messages, + request_kwargs, + pool_override=(routed_model,), + context_fit=context_fit, ) if pin_source_tier is not None else None @@ -3582,11 +3670,13 @@ class ComplexityRouter(CustomLogger): messages, resolved_messages, request_kwargs, + context_fit, ), messages, input, resolved_messages, request_kwargs, + context_fit, ) ) @@ -3598,14 +3688,18 @@ class ComplexityRouter(CustomLogger): specific_deployment=specific_deployment, conversation_continuing=conversation_continuing, resolved_messages=resolved_messages, + context_fit=context_fit, ) response: Final = ( await self._gate_response_health( - await self._gate_response_modality(routed_response, messages, resolved_messages, request_kwargs), + await self._gate_response_modality( + routed_response, messages, resolved_messages, request_kwargs, context_fit + ), messages, input, resolved_messages, request_kwargs, + context_fit, ) if routed_response is not None else None @@ -3640,6 +3734,7 @@ class ComplexityRouter(CustomLogger): specific_deployment: bool | None = False, conversation_continuing: bool = True, resolved_messages: Sequence[Mapping[str, object]] | None = None, + context_fit: _RequestContextFit | None = None, ) -> PreRoutingHookResponse | None: """ Classifies the request by complexity and returns the appropriate model. @@ -3811,7 +3906,9 @@ class ComplexityRouter(CustomLogger): plan_floored: Final = tier != pre_floor_tier if plan_floored: signals = (*signals, "plan_mode_floor") - context_placement: Final = await self._context_window_placement(tier, resolved_messages, request_kwargs) + context_placement: Final = await self._context_window_placement( + tier, resolved_messages, request_kwargs, context_fit=context_fit + ) tier, signals, context_original_tier = _apply_context_placement(tier, signals, context_placement) score_repr: Final = f"{score:.3f}" if score is not None else "n/a" fallback_model: Final = self.config.default_model if not self.config.plugins else None diff --git a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py index 6f3ea8eb78a..c7eb46046ef 100644 --- a/litellm/router_utils/pre_call_checks/deployment_affinity_check.py +++ b/litellm/router_utils/pre_call_checks/deployment_affinity_check.py @@ -349,7 +349,9 @@ class DeploymentAffinityCheck(CustomLogger): first write instead of the last. Re-claiming with the stored value refreshes its TTL, the same keepalive the complexity router's model pin documents: an active session must not lose its pin mid-conversation just because it outlives the - original write, so `session_affinity_ttl_seconds` bounds idle time, not total + original write, so the affinity TTL (the Router's + `deployment_affinity_ttl_seconds`, or a pre-routing hook's per-request + `session_affinity_ttl_seconds` override) bounds idle time, not total session length. On Redis one Lua script does the get-or-set-or-refresh atomically (same registration seam the rate limiters use) and the in-memory tier is synchronized to the winner; without Redis, and whenever Redis is diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index b7fdb5a98ef..89eab71ccba 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -2,13 +2,57 @@ from __future__ import annotations -from collections.abc import Awaitable +from collections.abc import Awaitable, Callable, Mapping, Sequence +from dataclasses import dataclass +from types import MappingProxyType from typing import Final, Protocol, cast # noqa: TID251 # native extension exposes dynamically typed callables import httpx -from litellm.rust_bridge.bindings import NativeBinding +import litellm +from litellm.constants import request_timeout +from litellm.llms.azure_ai.ocr.common_utils import is_azure_cohere_parse_model +from litellm.llms.base_llm.ocr.transformation import PROVIDER_NATIVE_RESPONSE_KEY, OCRResponse +from litellm.rust_bridge.bindings import NativeBinding, native_exception_types from litellm.rust_bridge.timeouts import timeout_to_seconds as _timeout_to_seconds +from litellm.types.router import GenericLiteLLMParams +from litellm.utils import ProviderConfigManager + +_RUST_OCR_PROVIDERS: Final = frozenset({"mistral", "azure_ai", "vertex_ai"}) +_RUST_OCR_CONFIG_FIELDS: Final = frozenset( + { + "azure_ad_token", + "tenant_id", + "client_id", + "client_secret", + "azure_scope", + "azure_authority_host", + "azure_credential", + "azure_federated_token_file", + "vertex_credentials", + "vertex_ai_credentials", + "vertex_project", + "vertex_ai_project", + "vertex_location", + "vertex_ai_location", + } +) +_RUST_OCR_SECRET_FIELDS: Final = frozenset( + {"azure_ad_token", "client_secret", "azure_federated_token_file", "vertex_credentials", "vertex_ai_credentials"} +) + + +@dataclass(frozen=True, slots=True) +class LiteLLMOcrRequest: + model: str + document: Mapping[str, object] + api_key: str | None + api_base: str | None + timeout: float | httpx.Timeout | None + custom_llm_provider: str | None + extra_headers: dict[str, object] | None + kwargs: Mapping[str, object] + input_sources: Mapping[str, str] | None = None class RustOcr(Protocol): @@ -21,6 +65,7 @@ class RustOcr(Protocol): custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], + input_sources: dict[str, str], timeout_seconds: float | None, ) -> dict[str, object]: raise NotImplementedError @@ -36,11 +81,32 @@ class RustAocr(Protocol): custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], + input_sources: dict[str, str], timeout_seconds: float | None, ) -> Awaitable[dict[str, object]]: raise NotImplementedError +class _OCRLogging(Protocol): + def update_from_kwargs( + self, + *, + kwargs: dict[str, object], + model: str, + optional_params: dict[str, object], + litellm_params: dict[str, object], + custom_llm_provider: str | None, + ) -> None: ... + + def pre_call( + self, + *, + input: str, + api_key: str | None, + additional_args: dict[str, object], + ) -> None: ... + + def _as_ocr(value: object) -> RustOcr | None: return cast(RustOcr, value) if callable(value) else None @@ -61,6 +127,264 @@ def load_rust_aocr() -> RustAocr | None: return _AOCR.load() +def provider(request: LiteLLMOcrRequest) -> str | None: + if request.custom_llm_provider is not None: + return request.custom_llm_provider + prefix: Final = request.model.partition("/")[0] + if prefix in _RUST_OCR_PROVIDERS: + return prefix + if request.model.startswith("mistral-ocr"): + return "mistral" + return None + + +def supported(request: LiteLLMOcrRequest) -> bool: + request_provider: Final = provider(request) + if request_provider not in _RUST_OCR_PROVIDERS: + return False + if request_provider == "azure_ai": + return ( + not is_azure_cohere_parse_model(request.model) + and not callable(request.kwargs.get("azure_ad_token_provider")) + and request.kwargs.get("azure_username") is None + and request.kwargs.get("azure_password") is None + ) + return True + + +def _optional_params(request: LiteLLMOcrRequest, resolve_secret: Callable[[str], str | None]) -> Mapping[str, object]: + optional_params: Final = MappingProxyType( + { + name: value + for name, value in request.kwargs.items() + if (name not in GenericLiteLLMParams.model_fields or name in _RUST_OCR_CONFIG_FIELDS) + and name not in ("litellm_logging_obj", "aocr", "litellm_call_id", "proxy_server_request") + } + ) + request_provider: Final = provider(request) + if request_provider == "azure_ai" and litellm.enable_azure_ad_token_refresh is True: + return MappingProxyType({**optional_params, "enable_azure_ad_token_refresh": True}) + if request_provider != "vertex_ai": + return optional_params + project: Final = ( + request.kwargs.get("vertex_project") + or request.kwargs.get("vertex_ai_project") + or litellm.vertex_project + or resolve_secret("VERTEXAI_PROJECT") + ) + location: Final = ( + request.kwargs.get("vertex_location") + or request.kwargs.get("vertex_ai_location") + or litellm.vertex_location + or resolve_secret("VERTEXAI_LOCATION") + or resolve_secret("VERTEX_LOCATION") + ) + credentials: Final = ( + request.kwargs.get("vertex_credentials") + or request.kwargs.get("vertex_ai_credentials") + or resolve_secret("VERTEXAI_CREDENTIALS") + ) + vertex_params: Final = MappingProxyType( + { + name: value + for name, value in ( + ("vertex_project", project), + ("vertex_location", location), + ("vertex_credentials", credentials), + ) + if value is not None + } + ) + return MappingProxyType({**optional_params, **vertex_params}) + + +def _input_sources(request: LiteLLMOcrRequest, optional_params: Mapping[str, object]) -> Mapping[str, str]: + proxy_request_value: Final = request.kwargs.get("proxy_server_request") + if not isinstance(proxy_request_value, Mapping): + return MappingProxyType({}) + proxy_request: Final = cast( # cast-ok: runtime Mapping check narrows metadata with unknown key and value types + Mapping[object, object], proxy_request_value + ) + credential_fields_value: Final = proxy_request.get("credential_fields", ()) + credential_fields: Final = ( + frozenset(name for name in credential_fields_value if isinstance(name, str)) + if isinstance(credential_fields_value, (list, tuple, set, frozenset)) + else frozenset() + ) + request_fields_value: Final = proxy_request.get("body_fields") + request_fields: Sequence[object] + if isinstance(request_fields_value, Sequence) and not isinstance(request_fields_value, (str, bytes)): + request_fields = cast( # cast-ok: runtime Sequence check excludes scalar strings and bytes + Sequence[object], request_fields_value + ) + else: + body_value: Final = proxy_request.get("body") + request_fields = ( + tuple(cast(Mapping[object, object], body_value)) # cast-ok: runtime Mapping check establishes iterable keys + if isinstance(body_value, Mapping) + else () + ) + names: Final = frozenset(optional_params) | frozenset({"api_key", "api_base", "extra_headers"}) + request_sources: Final = MappingProxyType( + {name: "request" for name in names if name in request_fields or name in credential_fields} + ) + if litellm.enable_azure_ad_token_refresh is True and "enable_azure_ad_token_refresh" in optional_params: + return MappingProxyType({**request_sources, "enable_azure_ad_token_refresh": "deployment"}) + return request_sources + + +def _marshal( + request: LiteLLMOcrRequest, + resolve_secret: Callable[[str], str | None], + convert_file_document: Callable[[dict[str, object]], dict[str, str]], +) -> LiteLLMOcrRequest: + if not isinstance(request.document, dict): + raise TypeError(f"document must be a dict with 'type' and URL/file field, got {type(request.document)}") + document: Final = ( + convert_file_document(request.document) if request.document.get("type") == "file" else request.document + ) + request_provider: Final = provider(request) + api_key: Final = ( + request.api_key or resolve_secret("MISTRAL_API_KEY") if request_provider == "mistral" else request.api_key + ) + optional_params: Final = _optional_params(request, resolve_secret) + input_sources: Final = _input_sources(request, optional_params) + logged_optional_params: Final = MappingProxyType( + {name: "****" if name in _RUST_OCR_SECRET_FIELDS else value for name, value in optional_params.items()} + ) + logged_kwargs: Final = MappingProxyType( + { + name: "****" if name in _RUST_OCR_SECRET_FIELDS else value + for name, value in request.kwargs.items() + if name != "proxy_server_request" + } + ) + logging_obj: Final = cast( # cast-ok: client decorator injects the logging object through untyped kwargs + _OCRLogging, request.kwargs["litellm_logging_obj"] + ) + logging_obj.update_from_kwargs( + kwargs=dict(logged_kwargs), # mutable-ok: legacy logging mutates its kwargs copy + model=request.model, + optional_params=dict(logged_optional_params), # mutable-ok: legacy logging requires concrete dict params + litellm_params={ # mutable-ok: legacy logging requires a concrete params dict + "litellm_call_id": request.kwargs.get("litellm_call_id"), + "api_base": request.api_base, + }, + custom_llm_provider=request_provider, + ) + logging_obj.pre_call( + input="OCR document processing", + api_key=api_key, + additional_args={ # mutable-ok: pre_call mutates the additional_args dict + "complete_input_dict": { # mutable-ok: callbacks consume a JSON-serializable request dict + "model": request.model, + "document": document, + **logged_optional_params, + }, + "api_base": request.api_base or "", + "headers": request.extra_headers or {}, # mutable-ok: logging callbacks consume a concrete headers dict + }, + ) + return LiteLLMOcrRequest( + model=request.model, + document=document, + api_key=api_key, + api_base=request.api_base, + timeout=request.timeout if request.timeout is not None else request_timeout, + custom_llm_provider=request.custom_llm_provider, + extra_headers=request.extra_headers, + kwargs=optional_params, + input_sources=input_sources, + ) + + +def _map_error(error: Exception, request: LiteLLMOcrRequest) -> Exception: + exception_types: Final = native_exception_types() + if exception_types is None or not isinstance(error, exception_types[1]): + return error + request_provider: Final = provider(request) + if request_provider is None: + return error + provider_config: Final = ProviderConfigManager.get_provider_ocr_config( + model=request.model.removeprefix(f"{request_provider}/"), provider=litellm.LlmProviders(request_provider) + ) + if provider_config is None: + return error + error_args: Final = cast( # cast-ok: BaseException.args exposes Any while native errors carry scalar args + tuple[object, ...], error.args + ) + status: Final = error_args[0] if error_args and isinstance(error_args[0], int) else 500 + message: Final = str(error_args[1]) if len(error_args) > 1 else str(error) + error_factory: Final = cast( # cast-ok: legacy provider error factories have untyped callable parameters + Callable[..., Exception], provider_config.get_error_class + ) + return error_factory( + error_message=message, + status_code=status or 500, + headers={}, # mutable-ok: provider error factories require a concrete headers dict + ) + + +def _response(response: Mapping[str, object]) -> OCRResponse: + provider_native_response: Final = response.get(PROVIDER_NATIVE_RESPONSE_KEY) + normalized: Final = OCRResponse.model_validate( + MappingProxyType({key: value for key, value in response.items() if key != PROVIDER_NATIVE_RESPONSE_KEY}) + ) + if isinstance(provider_native_response, Mapping): + normalized.set_provider_native_response(provider_native_response) + return normalized + + +def run( + request: LiteLLMOcrRequest, + resolve_secret: Callable[[str], str | None], + convert_file_document: Callable[[dict[str, object]], dict[str, str]], +) -> OCRResponse | None: + if load_rust_ocr() is None: + return None + marshalled: Final = _marshal(request, resolve_secret, convert_file_document) + try: + response: Final = ocr( + model=marshalled.model, + document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict + api_key=marshalled.api_key, + api_base=marshalled.api_base, + custom_llm_provider=marshalled.custom_llm_provider, + extra_headers=marshalled.extra_headers, + optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict + input_sources=marshalled.input_sources, + timeout=marshalled.timeout, + ) + except Exception as error: + raise _map_error(error, request) from error + return _response(response) if response is not None else None + + +async def arun( + request: LiteLLMOcrRequest, + resolve_secret: Callable[[str], str | None], + convert_file_document: Callable[[dict[str, object]], dict[str, str]], +) -> OCRResponse | None: + if load_rust_aocr() is None: + return None + marshalled: Final = _marshal(request, resolve_secret, convert_file_document) + try: + response: Final = await aocr( + model=marshalled.model, + document=dict(marshalled.document), # mutable-ok: PyO3 OCR binding requires a concrete dict + api_key=marshalled.api_key, + api_base=marshalled.api_base, + custom_llm_provider=marshalled.custom_llm_provider, + extra_headers=marshalled.extra_headers, + optional_params=dict(marshalled.kwargs), # mutable-ok: PyO3 OCR binding requires a concrete dict + input_sources=marshalled.input_sources, + timeout=marshalled.timeout, + ) + except Exception as error: + raise _map_error(error, request) from error + return _response(response) if response is not None else None + + def ocr( *, model: str, @@ -71,6 +395,7 @@ def ocr( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + input_sources: Mapping[str, str] | None = None, ) -> dict[str, object] | None: rust_ocr: Final = load_rust_ocr() if rust_ocr is None: @@ -83,6 +408,7 @@ def ocr( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, optional_params=optional_params, + input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict timeout_seconds=_timeout_to_seconds(timeout), ) @@ -97,6 +423,7 @@ async def aocr( extra_headers: dict[str, object] | None, optional_params: dict[str, object], timeout: float | httpx.Timeout | None, + input_sources: Mapping[str, str] | None = None, ) -> dict[str, object] | None: rust_aocr: Final = load_rust_aocr() if rust_aocr is None: @@ -109,5 +436,6 @@ async def aocr( custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, optional_params=optional_params, + input_sources=dict(input_sources or {}), # mutable-ok: native boundary requires a concrete dict timeout_seconds=_timeout_to_seconds(timeout), ) diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index e86c8e7c919..d75375a01cc 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -47,6 +47,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): aws_web_identity_token: str | None = None, aws_sts_endpoint: str | None = None, replica_regions: list[str] | None = None, + kms_key_id: str | None = None, **kwargs, ): BaseSecretManager.__init__(self, **kwargs) @@ -61,6 +62,7 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): self.aws_web_identity_token = aws_web_identity_token self.aws_sts_endpoint = aws_sts_endpoint self.replica_regions: list[str] = replica_regions or [] + self.kms_key_id = kms_key_id @classmethod def validate_environment(cls): @@ -106,7 +108,8 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): # Remove None values aws_kwargs = {k: v for k, v in aws_kwargs.items() if v is not None} - litellm.secret_manager_client = cls(**aws_kwargs) + kms_key_id: Final = key_management_settings.kms_key_id if key_management_settings is not None else None + litellm.secret_manager_client = cls(kms_key_id=kms_key_id, **aws_kwargs) litellm._key_management_system = KeyManagementSystem.AWS_SECRET_MANAGER except Exception as e: @@ -275,6 +278,9 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): if description: data["Description"] = description + if self.kms_key_id: + data["KmsKeyId"] = self.kms_key_id + # ✅ Normalize tags to AWS format if tags: if isinstance(tags, dict): diff --git a/litellm/types/integrations/datadog_llm_obs.py b/litellm/types/integrations/datadog_llm_obs.py index 17cf5831c96..f4fcabf53ed 100644 --- a/litellm/types/integrations/datadog_llm_obs.py +++ b/litellm/types/integrations/datadog_llm_obs.py @@ -87,6 +87,7 @@ class LLMMetrics(TypedDict, total=False): cache_write_input_tokens: ReadOnly[float] non_cached_input_tokens: ReadOnly[float] reasoning_output_tokens: ReadOnly[float] + tool_output_tokens: ReadOnly[float] class LLMObsPayload(TypedDict, total=False): diff --git a/litellm/types/management_endpoints/auto_router_endpoints.py b/litellm/types/management_endpoints/auto_router_endpoints.py index 6306658ad0b..9f29f27e41d 100644 --- a/litellm/types/management_endpoints/auto_router_endpoints.py +++ b/litellm/types/management_endpoints/auto_router_endpoints.py @@ -226,6 +226,30 @@ class AutoRouterBenchmarkGroup(AutoRouterBenchmarkTotals): ) +class AutoRouterSessionResponse(BaseModel): + """One auto-routed session as its own key sees it: what the last turn ran on, and what the session cost + against the router's savings baseline (the priciest model in its hardest tier).""" + + session_id: str + router_name: str = Field(description="The auto-router alias the session's requests were sent to") + router_type: str = Field(description="complexity, adaptive or quality") + turns: int = Field(description="Auto-routed turns the rollup has recorded for this session so far") + last_model: str = Field(description="The deployment model the most recent turn was routed to") + spend: float = Field(description="What the session's routed traffic actually cost, classifier calls included") + saved_spend: float = Field(description="Estimated savings against the baseline, net of classifier cost") + baseline_spend: float = Field(description="spend plus saved_spend: the estimated single-model cost") + baseline_model: str | None = Field( + description="The savings baseline most of this session's turns were priced against, recorded turn by " + "turn, so it still names the counterfactual after the router is reconfigured or removed. None when no " + "turn recorded one: rows from before the baseline was recorded, and adaptive and quality routers, " + "which derive no baseline and so report no savings" + ) + baseline_models: Mapping[str, int] = Field( + description="Turns priced against each baseline model; more than one entry means the router's " + "baseline changed mid-session and baseline_spend mixes both" + ) + + class AutoRouterBenchmarksResponse(BaseModel): """Benchmarks for the auto-router dashboard, aggregated from the per-session rollup.""" diff --git a/litellm/types/secret_managers/main.py b/litellm/types/secret_managers/main.py index 599e5746dfb..148e680a236 100644 --- a/litellm/types/secret_managers/main.py +++ b/litellm/types/secret_managers/main.py @@ -45,6 +45,9 @@ class KeyManagementSettings(LiteLLMPydanticObjectBase): tags: dict[str, str] | None = None """Optional tags to attach when creating secrets (e.g. {"Environment": "Prod", "Owner": "AI-Platform"}).""" + kms_key_id: str | None = None + """Optional customer-managed KMS key (ID, alias or ARN) used to encrypt secrets created in AWS Secrets Manager.""" + custom_secret_manager: str | None = None """ Path to custom secret manager class (e.g. "my_secret_manager.InMemorySecretManager") diff --git a/litellm/types/utils.py b/litellm/types/utils.py index 58b940227f8..e3ea37dc0c8 100644 --- a/litellm/types/utils.py +++ b/litellm/types/utils.py @@ -2920,6 +2920,7 @@ RoutingDecisionCause = Literal[ # same tier served instead. The displaced group rides in signals. Reported even on a kept # session pin, since the pinned model did not serve the request. "health_failover", + "health_default_fallback", "session_affinity_pin", "session_affinity_escalation", # classification_mode 'user_turn': the request is an agent loop's continuation turn (no new diff --git a/litellm/types/vector_stores.py b/litellm/types/vector_stores.py index 8bb0235ea2a..6d2ca308798 100644 --- a/litellm/types/vector_stores.py +++ b/litellm/types/vector_stores.py @@ -5,7 +5,7 @@ from enum import Enum from typing import Any, Literal from pydantic import BaseModel -from typing_extensions import TypedDict +from typing_extensions import ReadOnly, TypedDict class SupportedVectorStoreIntegrations(str, Enum): @@ -96,6 +96,17 @@ class VectorStoreSearchResponse(TypedDict, total=False): data: list[VectorStoreSearchResult] | None +VectorStoreSearchFailureMode = Literal["annotate", "error"] + + +class VectorStoreSearchFailure(TypedDict): + """A configured vector store whose search failed, as reported back to the API caller""" + + vector_store_id: ReadOnly[str] + custom_llm_provider: ReadOnly[str | None] + error: ReadOnly[str] + + class VectorStoreSearchOptionalRequestParams(TypedDict, total=False): """TypedDict for Optional parameters supported by the vector store search API.""" diff --git a/litellm/utils.py b/litellm/utils.py index a765e1b1246..bb1bce66d9b 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -9400,11 +9400,9 @@ class ProviderConfigManager: ReductoParseV3Config, ) - if model == "parse-v3": - return ReductoParseV3Config() if model == "parse-legacy": return ReductoParseLegacyConfig() - return None + return ReductoParseV3Config() MistralOCRConfig: Final = litellm_utils.MistralOCRConfig PROVIDER_TO_CONFIG_MAP: Final = { diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b7726290f0e..aadd9bd3028 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -13118,6 +13118,22 @@ "supports_reasoning": true, "supports_tool_choice": true }, + "cerebras/qwen-3.8-27b": { + "input_cost_per_token": 9.9e-07, + "litellm_provider": "cerebras", + "max_input_tokens": 65536, + "max_output_tokens": 32768, + "max_tokens": 32768, + "mode": "chat", + "output_cost_per_token": 1.49e-06, + "source": "https://api.cerebras.ai/public/v1/models/qwen-3.8-27b", + "supports_function_calling": true, + "supports_parallel_function_calling": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_vision": true + }, "chatdolphin": { "input_cost_per_token": 5e-07, "litellm_provider": "nlp_cloud", @@ -31328,8 +31344,6 @@ "cache_read_input_token_cost_above_272k_tokens_flex": 5e-07 }, "gpt-5.5-pro": { - "cache_read_input_token_cost": 3e-06, - "cache_read_input_token_cost_above_272k_tokens": 6e-06, "input_cost_per_token": 3e-05, "input_cost_per_token_above_272k_tokens": 6e-05, "input_cost_per_token_flex": 1.5e-05, @@ -31378,8 +31392,6 @@ "supports_low_reasoning_effort": false }, "gpt-5.5-pro-2026-04-23": { - "cache_read_input_token_cost": 3e-06, - "cache_read_input_token_cost_above_272k_tokens": 6e-06, "input_cost_per_token": 3e-05, "input_cost_per_token_above_272k_tokens": 6e-05, "input_cost_per_token_flex": 1.5e-05, @@ -31532,8 +31544,6 @@ "cache_read_input_token_cost_above_272k_tokens_flex": 2.5e-07 }, "gpt-5.4-pro": { - "cache_read_input_token_cost": 3e-06, - "cache_read_input_token_cost_above_272k_tokens": 6e-06, "input_cost_per_token": 3e-05, "input_cost_per_token_above_272k_tokens": 6e-05, "input_cost_per_token_flex": 1.5e-05, @@ -31583,8 +31593,6 @@ "output_cost_per_token_above_272k_tokens_flex": 0.000135 }, "gpt-5.4-pro-2026-03-05": { - "cache_read_input_token_cost": 3e-06, - "cache_read_input_token_cost_above_272k_tokens": 6e-06, "input_cost_per_token": 3e-05, "input_cost_per_token_above_272k_tokens": 6e-05, "input_cost_per_token_flex": 1.5e-05, @@ -33076,6 +33084,7 @@ "supports_tool_choice": true }, "groq/gemma-7b-it": { + "deprecation_date": "2024-12-18", "input_cost_per_token": 5e-08, "litellm_provider": "groq", "max_input_tokens": 8192, @@ -33890,6 +33899,20 @@ "supports_system_messages": true, "supports_tool_choice": true }, + "inception/mercury-2.5": { + "input_cost_per_token": 2e-07, + "litellm_provider": "inception", + "max_input_tokens": 260000, + "max_output_tokens": 65536, + "max_tokens": 65536, + "mode": "chat", + "output_cost_per_token": 7.5e-07, + "source": "https://docs.inceptionlabs.ai/get-started/models", + "supports_function_calling": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true + }, "text-completion-inception/mercury-edit-2": { "cache_read_input_token_cost": 2.5e-08, "input_cost_per_token": 2.5e-07, @@ -39126,19 +39149,21 @@ "supports_tool_choice": true }, "openrouter/deepseek/deepseek-chat-v3.1": { - "input_cost_per_token": 2e-07, + "input_cost_per_token": 2.5e-07, "input_cost_per_token_cache_hit": 2e-08, "litellm_provider": "openrouter", "max_input_tokens": 163840, "max_output_tokens": 163840, "max_tokens": 163840, "mode": "chat", - "output_cost_per_token": 8e-07, + "output_cost_per_token": 9.5e-07, "supports_assistant_prefill": true, "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "cache_read_input_token_cost": 1.3e-07, + "source": "https://openrouter.ai/deepseek/deepseek-chat-v3.1" }, "openrouter/deepseek/deepseek-v3.2": { "input_cost_per_token": 2.69e-07, @@ -39201,36 +39226,56 @@ "supports_tool_choice": true }, "openrouter/deepseek/deepseek-v4-pro": { - "input_cost_per_token": 1.32e-06, + "input_cost_per_token": 8.59908e-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": 3.96e-06, + "output_cost_per_token": 1.719816e-06, "source": "https://openrouter.ai/deepseek/deepseek-v4-pro", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "cache_read_input_token_cost": 7.1659e-08 + }, + "openrouter/deepseek/deepseek-v4.1-flash": { + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 6e-07, + "cache_read_input_token_cost": 3e-09, + "litellm_provider": "openrouter", + "max_input_tokens": 1048576, + "max_output_tokens": 384000, + "max_tokens": 384000, + "mode": "chat", + "source": "https://openrouter.ai/deepseek/deepseek-v4.1-flash", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": false, + "supports_prompt_caching": true }, "openrouter/deepseek/deepseek-v4-pro-0813": { - "input_cost_per_token": 1.32e-06, + "input_cost_per_token": 5.7948e-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": 3.96e-06, + "output_cost_per_token": 1.73844e-06, "source": "https://openrouter.ai/deepseek/deepseek-v4-pro-0813", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "cache_read_input_token_cost": 1.9316e-08 }, "openrouter/google/gemini-2.0-flash-001": { "deprecation_date": "2026-06-01", @@ -40009,6 +40054,28 @@ "supports_tool_choice": true, "supports_vision": true }, + "openrouter/openai/gpt-5.6-sol-pro": { + "input_cost_per_token": 2e-06, + "output_cost_per_token": 1e-05, + "cache_read_input_token_cost": 2e-07, + "cache_creation_input_token_cost": 2.5e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-5.6-sol-pro", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/openai/gpt-oss-120b": { "input_cost_per_token": 3.7e-08, "litellm_provider": "openrouter", @@ -40137,13 +40204,13 @@ "supports_tool_choice": true }, "openrouter/qwen/qwen3-235b-a22b-2507": { - "input_cost_per_token": 8.75e-08, + "input_cost_per_token": 2.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 262144, "max_tokens": 262144, "mode": "chat", - "output_cost_per_token": 3.5e-07, + "output_cost_per_token": 8.8e-07, "source": "https://openrouter.ai/qwen/qwen3-235b-a22b-2507", "supports_function_calling": true, "supports_tool_choice": true @@ -40176,7 +40243,7 @@ "supports_vision": true }, "openrouter/qwen/qwen3.5-35b-a3b": { - "input_cost_per_token": 2.5e-07, + "input_cost_per_token": 3.125e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, @@ -40187,7 +40254,8 @@ "supports_function_calling": true, "supports_reasoning": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "cache_read_input_token_cost": 1.5625e-07 }, "openrouter/qwen/qwen3.5-27b": { "input_cost_per_token": 1.95e-07, @@ -40204,13 +40272,13 @@ "supports_vision": true }, "openrouter/qwen/qwen3.5-122b-a10b": { - "input_cost_per_token": 2.9e-07, + "input_cost_per_token": 2.6e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 65536, "max_tokens": 65536, "mode": "chat", - "output_cost_per_token": 2.4e-06, + "output_cost_per_token": 2.08e-06, "source": "https://openrouter.ai/qwen/qwen3.5-122b-a10b", "supports_function_calling": true, "supports_reasoning": true, @@ -40297,18 +40365,19 @@ "supports_web_search": true }, "openrouter/z-ai/glm-4.6": { - "input_cost_per_token": 5.5e-07, + "input_cost_per_token": 4.3e-07, "litellm_provider": "openrouter", "max_input_tokens": 202800, "max_output_tokens": 131000, "max_tokens": 131000, "mode": "chat", - "output_cost_per_token": 2.2e-06, + "output_cost_per_token": 1.75e-06, "source": "https://openrouter.ai/z-ai/glm-4.6", "supports_function_calling": true, "supports_prompt_caching": true, "supports_reasoning": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "cache_read_input_token_cost": 8e-08 }, "openrouter/z-ai/glm-4.6:exacto": { "input_cost_per_token": 4.5e-07, @@ -43133,7 +43202,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 1.2e-06, + "max_input_tokens": 131072, + "source": "https://api.together.xyz/v1/models" }, "together_ai/Qwen/Qwen2.5-7B-Instruct-Turbo": { "litellm_provider": "together_ai", @@ -43141,7 +43214,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "input_cost_per_token": 3e-07, + "output_cost_per_token": 3e-07, + "max_input_tokens": 32768, + "source": "https://api.together.xyz/v1/models" }, "together_ai/Qwen/Qwen3-235B-A22B-Instruct-2507-tput": { "deprecation_date": "2026-07-10", @@ -43353,7 +43430,11 @@ "supports_function_calling": true, "supports_parallel_function_calling": true, "supports_response_schema": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "input_cost_per_token": 2e-07, + "output_cost_per_token": 2e-07, + "max_input_tokens": 32768, + "source": "https://api.together.xyz/v1/models" }, "together_ai/mistralai/Mistral-Small-24B-Instruct-2501": { "deprecation_date": "2026-04-02", @@ -43361,7 +43442,11 @@ "mode": "chat", "supports_function_calling": true, "supports_parallel_function_calling": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "input_cost_per_token": 1e-07, + "output_cost_per_token": 3e-07, + "max_input_tokens": 32768, + "source": "https://api.together.xyz/v1/models" }, "together_ai/mistralai/Mixtral-8x7B-Instruct-v0.1": { "deprecation_date": "2026-04-16", @@ -48092,27 +48177,29 @@ "supports_tool_choice": true }, "vertex_ai/mistral-small-2503": { - "input_cost_per_token": 1e-06, + "input_cost_per_token": 1e-07, "litellm_provider": "vertex_ai-mistral_models", "max_input_tokens": 128000, "max_output_tokens": 128000, "max_tokens": 128000, "mode": "chat", - "output_cost_per_token": 3e-06, + "output_cost_per_token": 3e-07, "supports_function_calling": true, "supports_tool_choice": true, - "supports_vision": true + "supports_vision": true, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing" }, "vertex_ai/mistral-small-2503@001": { - "input_cost_per_token": 1e-06, + "input_cost_per_token": 1e-07, "litellm_provider": "vertex_ai-mistral_models", "max_input_tokens": 32000, "max_output_tokens": 8191, "max_tokens": 8191, "mode": "chat", - "output_cost_per_token": 3e-06, + "output_cost_per_token": 3e-07, "supports_function_calling": true, - "supports_tool_choice": true + "supports_tool_choice": true, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing" }, "vertex_ai/mistral-ocr-2505": { "litellm_provider": "vertex_ai", @@ -48162,15 +48249,16 @@ "supports_reasoning": true }, "vertex_ai/openai/gpt-oss-20b-maas": { - "input_cost_per_token": 7.5e-08, + "input_cost_per_token": 7e-08, "litellm_provider": "vertex_ai-openai_models", "max_input_tokens": 131072, "max_output_tokens": 32768, "max_tokens": 32768, "mode": "chat", - "output_cost_per_token": 3e-07, - "source": "https://console.cloud.google.com/vertex-ai/publishers/openai/model-garden/gpt-oss-120b-maas", - "supports_reasoning": true + "output_cost_per_token": 2.5e-07, + "source": "https://cloud.google.com/vertex-ai/generative-ai/pricing", + "supports_reasoning": true, + "cache_read_input_token_cost": 7e-09 }, "vertex_ai/xai/grok-4.1-fast-non-reasoning": { "cache_read_input_token_cost": 5e-08, @@ -55272,11 +55360,14 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "bedrock_mantle/openai.gpt-5.6-terra": { @@ -55311,11 +55402,14 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "bedrock_mantle/openai.gpt-5.6-cyber": { @@ -55340,10 +55434,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "bedrock_mantle/openai.gpt-daybreak-blue-5.6-sol": { @@ -55372,10 +55468,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-daybreak-blue-56-sol.html" }, @@ -55411,11 +55509,14 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "us.openai.gpt-5.6-sol": { @@ -55440,8 +55541,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "global.openai.gpt-5.6-sol": { @@ -55466,8 +55570,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "us.openai.gpt-5.6-terra": { @@ -55492,8 +55599,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "global.openai.gpt-5.6-terra": { @@ -55518,8 +55628,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "us.openai.gpt-5.6-luna": { @@ -55544,8 +55657,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "global.openai.gpt-5.6-luna": { @@ -55570,8 +55686,11 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, "supports_tool_choice": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true }, "bedrock_mantle/openai.gpt-6-astra": { @@ -55601,10 +55720,14 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-astra.html" }, @@ -55630,9 +55753,13 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, "supports_tool_choice": true, "supports_prompt_caching": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-astra.html" }, @@ -55658,9 +55785,13 @@ "text" ], "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, "supports_tool_choice": true, "supports_prompt_caching": true, "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-astra.html" }, @@ -55693,11 +55824,13 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "bedrock_mantle/openai.gpt-5.4": { @@ -55729,11 +55862,13 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, "supports_vision": true, + "supports_xhigh_reasoning_effort": true, "supports_web_search": true }, "bedrock_mantle/google.gemma-4-31b": { @@ -56465,17 +56600,17 @@ "supports_reasoning": true, "source": "https://serverless.tensormesh.ai/v1/models/openrouter" }, - "deepseek-v4-flash": { + "deepseek-flash": { "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 1.4e-08, - "input_cost_per_token": 4.4e-07, - "input_cost_per_token_cache_hit": 1.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, "litellm_provider": "deepseek", "max_input_tokens": 1000000, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.32e-06, + "output_cost_per_token": 1.2e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", "supported_endpoints": [ "/v1/chat/completions" @@ -56489,19 +56624,45 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": false + "supports_vision": true }, - "deepseek-v4-flash-vision-exp": { + "deepseek-v4-flash": { "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 1.4e-08, - "input_cost_per_token": 4.4e-07, - "input_cost_per_token_cache_hit": 1.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, "litellm_provider": "deepseek", "max_input_tokens": 1000000, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.32e-06, + "output_cost_per_token": 1.2e-06, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "deepseek-v4-flash-vision-exp": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 393216, + "max_tokens": 393216, + "mode": "chat", + "output_cost_per_token": 1.2e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", "supported_endpoints": [ "/v1/chat/completions" @@ -56543,17 +56704,17 @@ "supports_tool_choice": true, "supports_vision": false }, - "deepseek/deepseek-v4-flash": { + "deepseek/deepseek-flash": { "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 1.4e-08, - "input_cost_per_token": 4.4e-07, - "input_cost_per_token_cache_hit": 1.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, "litellm_provider": "deepseek", "max_input_tokens": 1000000, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.32e-06, + "output_cost_per_token": 1.2e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", "supported_endpoints": [ "/v1/chat/completions" @@ -56567,19 +56728,45 @@ "supports_response_schema": true, "supports_system_messages": true, "supports_tool_choice": true, - "supports_vision": false + "supports_vision": true }, - "deepseek/deepseek-v4-flash-vision-exp": { + "deepseek/deepseek-v4-flash": { "cache_creation_input_token_cost": 0.0, - "cache_read_input_token_cost": 1.4e-08, - "input_cost_per_token": 4.4e-07, - "input_cost_per_token_cache_hit": 1.4e-08, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, "litellm_provider": "deepseek", "max_input_tokens": 1000000, "max_output_tokens": 393216, "max_tokens": 393216, "mode": "chat", - "output_cost_per_token": 1.32e-06, + "output_cost_per_token": 1.2e-06, + "source": "https://api-docs.deepseek.com/quick_start/pricing", + "supported_endpoints": [ + "/v1/chat/completions" + ], + "supports_assistant_prefill": true, + "supports_function_calling": true, + "supports_native_streaming": true, + "supports_parallel_function_calling": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true + }, + "deepseek/deepseek-v4-flash-vision-exp": { + "cache_creation_input_token_cost": 0.0, + "cache_read_input_token_cost": 6e-09, + "input_cost_per_token": 3e-07, + "input_cost_per_token_cache_hit": 6e-09, + "litellm_provider": "deepseek", + "max_input_tokens": 1000000, + "max_output_tokens": 393216, + "max_tokens": 393216, + "mode": "chat", + "output_cost_per_token": 1.2e-06, "source": "https://api-docs.deepseek.com/quick_start/pricing", "supported_endpoints": [ "/v1/chat/completions" @@ -56935,6 +57122,23 @@ ], "supports_audio_input": true }, + "gpt-live-1": { + "input_cost_per_second": 0.0008333333333333334, + "litellm_provider": "openai", + "mode": "realtime", + "source": "https://developers.openai.com/api/docs/models/gpt-live-1", + "supported_modalities": [ + "text", + "audio" + ], + "supported_output_modalities": [ + "text", + "audio" + ], + "supports_audio_input": true, + "supports_audio_output": true, + "supports_function_calling": true + }, "gpt-realtime-translate": { "input_cost_per_second": 0.0005666666666666667, "litellm_provider": "openai", @@ -60039,6 +60243,7 @@ ] }, "xai/grok-imagine-image-quality": { + "deprecation_date": "2026-11-02", "input_cost_per_image": 0.05, "litellm_provider": "xai", "mode": "image_generation", @@ -60055,6 +60260,7 @@ ] }, "xai/grok-imagine-image-quality-20260403": { + "deprecation_date": "2026-11-02", "input_cost_per_image": 0.05, "litellm_provider": "xai", "mode": "image_generation", @@ -60071,6 +60277,7 @@ ] }, "xai/grok-imagine-image-quality-latest": { + "deprecation_date": "2026-11-02", "input_cost_per_image": 0.05, "litellm_provider": "xai", "mode": "image_generation", @@ -60590,6 +60797,113 @@ "output_cost_per_token": 4.7e-07, "source": "https://docs.together.ai/docs/serverless-models" }, + "together_ai/moonshotai/Kimi-K2.6": { + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 4.5e-06, + "cache_read_input_token_cost": 2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/moonshotai/Kimi-K2.5-fp4": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 2.8e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/MiniMaxAI/MiniMax-M2.7": { + "input_cost_per_token": 3e-07, + "output_cost_per_token": 1.2e-06, + "cache_read_input_token_cost": 6e-08, + "litellm_provider": "together_ai", + "max_input_tokens": 196608, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/zai-org/GLM-5": { + "input_cost_per_token": 1e-06, + "output_cost_per_token": 3.2e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 202752, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/zai-org/GLM-5.1": { + "input_cost_per_token": 1.4e-06, + "output_cost_per_token": 4.4e-06, + "cache_read_input_token_cost": 2.6e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 202752, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/deepseek-ai/DeepSeek-R1-0528": { + "input_cost_per_token": 3e-06, + "output_cost_per_token": 7e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 163840, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/Qwen/Qwen3-Coder-Next-FP8": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 1.2e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/Qwen/Qwen3-VL-32B-Instruct": { + "input_cost_per_token": 5e-07, + "output_cost_per_token": 1.5e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/Qwen/Qwen3-VL-8B-Instruct": { + "input_cost_per_token": 1.8e-07, + "output_cost_per_token": 6.8e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/mistralai/Ministral-3-14B-Instruct-2512": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 262144, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/nvidia/NVIDIA-Nemotron-Nano-9B-v2": { + "input_cost_per_token": 6e-08, + "output_cost_per_token": 2.5e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 131072, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/mistralai/Mistral-7B-Instruct-v0.3": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 2e-07, + "litellm_provider": "together_ai", + "max_input_tokens": 32768, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, + "together_ai/Qwen/QwQ-32B": { + "input_cost_per_token": 1.2e-06, + "output_cost_per_token": 1.2e-06, + "litellm_provider": "together_ai", + "max_input_tokens": 131072, + "mode": "chat", + "source": "https://api.together.xyz/v1/models" + }, "cerebras/gemma-4-31b": { "input_cost_per_token": 9.9e-07, "litellm_provider": "cerebras", @@ -61093,10 +61407,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "input_cost_per_token": 2.64e-06, "input_cost_per_token_above_272k_tokens": 5.28e-06, @@ -61126,10 +61442,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "input_cost_per_token": 2.64e-07, "input_cost_per_token_above_272k_tokens": 5.28e-07, @@ -61158,10 +61476,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "input_cost_per_token": 3.3e-06, "cache_read_input_token_cost": 3.3e-07, @@ -61320,10 +61640,12 @@ "text" ], "supports_function_calling": true, + "supports_minimal_reasoning_effort": false, "supports_prompt_caching": true, "supports_reasoning": true, "supports_response_schema": true, "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, "supports_vision": true, "input_cost_per_token": 3.3e-06, "cache_read_input_token_cost": 3.3e-07, @@ -61860,6 +62182,7 @@ "mode": "responses", "supports_web_search": true, "supports_function_calling": true, + "supports_reasoning": true, "input_cost_per_token": 1.15e-08, "output_cost_per_token": 1.7e-07, "cache_read_input_token_cost": 1.15e-09, @@ -61870,6 +62193,7 @@ "mode": "responses", "supports_web_search": true, "supports_function_calling": true, + "supports_reasoning": true, "input_cost_per_token": 2.5e-07, "output_cost_per_token": 2.5e-06, "cache_read_input_token_cost": 2.5e-07, @@ -62233,6 +62557,28 @@ "cache_read_input_token_cost": 2e-08, "supports_prompt_caching": true }, + "openrouter/openai/gpt-5.6-luna-pro": { + "input_cost_per_token": 2e-07, + "output_cost_per_token": 1.2e-06, + "cache_read_input_token_cost": 2e-08, + "cache_creation_input_token_cost": 2.5e-07, + "input_cost_per_token_above_272k_tokens": 4e-07, + "output_cost_per_token_above_272k_tokens": 1.8e-06, + "cache_read_input_token_cost_above_272k_tokens": 4e-08, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-5.6-luna-pro", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/openai/gpt-5.6-terra": { "input_cost_per_token": 2e-06, "output_cost_per_token": 1.2e-05, @@ -62252,6 +62598,28 @@ "cache_read_input_token_cost": 2e-07, "supports_prompt_caching": true }, + "openrouter/openai/gpt-5.6-terra-pro": { + "input_cost_per_token": 2e-06, + "output_cost_per_token": 1.2e-05, + "cache_read_input_token_cost": 2e-07, + "cache_creation_input_token_cost": 2.5e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "output_cost_per_token_above_272k_tokens": 1.8e-05, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-5.6-terra-pro", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/openai/o3": { "input_cost_per_token": 2e-06, "output_cost_per_token": 8e-06, @@ -62485,6 +62853,28 @@ "supports_pdf_input": true, "supports_prompt_caching": true }, + "openrouter/openai/gpt-6-astra-pro": { + "input_cost_per_token": 1e-05, + "output_cost_per_token": 5e-05, + "cache_read_input_token_cost": 1e-06, + "cache_creation_input_token_cost": 1.25e-05, + "input_cost_per_token_above_272k_tokens": 2e-05, + "output_cost_per_token_above_272k_tokens": 7.5e-05, + "cache_read_input_token_cost_above_272k_tokens": 2e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-6-astra-pro", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/qwen/qwen3.8-flash": { "input_cost_per_token": 1.5e-07, "output_cost_per_token": 4.7e-07, @@ -62504,9 +62894,9 @@ "supports_prompt_caching": true }, "openrouter/z-ai/glm-5.3-flash": { - "input_cost_per_token": 7.5e-08, - "output_cost_per_token": 2.5e-07, - "cache_read_input_token_cost": 1.5e-08, + "input_cost_per_token": 1.5e-07, + "output_cost_per_token": 5e-07, + "cache_read_input_token_cost": 3e-08, "litellm_provider": "openrouter", "max_input_tokens": 1310720, "max_output_tokens": 131072, @@ -62621,6 +63011,25 @@ "supports_vision": true, "supports_prompt_caching": true }, + "openrouter/qwen/qwen3.8-max-0902": { + "input_cost_per_token": 2e-06, + "output_cost_per_token": 6e-06, + "cache_read_input_token_cost": 2.5e-07, + "cache_creation_input_token_cost": 2.5e-06, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 131072, + "max_tokens": 131072, + "mode": "chat", + "source": "https://openrouter.ai/qwen/qwen3.8-max-0902", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": false, + "supports_prompt_caching": true + }, "openrouter/deepseek/deepseek-v4-flash-0731": { "input_cost_per_token": 6.5e-08, "output_cost_per_token": 1.8e-07, @@ -62692,9 +63101,9 @@ "supports_vision": false }, "openrouter/moonshotai/kimi-k3": { - "input_cost_per_token": 3e-06, - "output_cost_per_token": 1.5e-05, - "cache_read_input_token_cost": 3e-07, + "input_cost_per_token": 2.1e-06, + "output_cost_per_token": 1.053e-05, + "cache_read_input_token_cost": 2.35e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 943718, @@ -62791,9 +63200,9 @@ "supports_prompt_caching": true }, "openrouter/z-ai/glm-5.2": { - "input_cost_per_token": 9.66e-07, - "output_cost_per_token": 3.036e-06, - "cache_read_input_token_cost": 1.932e-07, + "input_cost_per_token": 6e-07, + "output_cost_per_token": 2e-06, + "cache_read_input_token_cost": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 131072, @@ -62824,9 +63233,9 @@ "supports_vision": false }, "openrouter/moonshotai/kimi-k2.7-code": { - "input_cost_per_token": 6.6e-07, - "output_cost_per_token": 3.4e-06, - "cache_read_input_token_cost": 1.8e-07, + "input_cost_per_token": 7.1e-07, + "output_cost_per_token": 3.5e-06, + "cache_read_input_token_cost": 1.5e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 235929, @@ -63074,10 +63483,28 @@ "supports_vision": true, "supports_pdf_input": true }, + "openrouter/openai/gpt-chat-latest": { + "input_cost_per_token": 5e-06, + "output_cost_per_token": 3e-05, + "cache_read_input_token_cost": 5e-07, + "litellm_provider": "openrouter", + "max_input_tokens": 400000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "source": "https://openrouter.ai/openai/gpt-chat-latest", + "supports_function_calling": true, + "supports_tool_choice": true, + "supports_reasoning": false, + "supports_response_schema": true, + "supports_vision": true, + "supports_pdf_input": true, + "supports_prompt_caching": true + }, "openrouter/deepseek/deepseek-v4-flash": { - "input_cost_per_token": 8.778e-08, - "output_cost_per_token": 1.7556e-07, - "cache_read_input_token_cost": 1.7556e-08, + "input_cost_per_token": 8.54e-08, + "output_cost_per_token": 1.708e-07, + "cache_read_input_token_cost": 1.708e-08, "litellm_provider": "openrouter", "max_input_tokens": 1048576, "max_output_tokens": 384000, @@ -63110,8 +63537,8 @@ "supports_prompt_caching": true }, "openrouter/google/gemma-4-26b-a4b-it": { - "input_cost_per_token": 7e-08, - "output_cost_per_token": 3.4e-07, + "input_cost_per_token": 4.2e-08, + "output_cost_per_token": 2.2e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 16384, @@ -63793,7 +64220,7 @@ "supports_vision": false }, "openrouter/qwen/qwen3-next-80b-a3b-instruct": { - "input_cost_per_token": 1e-07, + "input_cost_per_token": 9e-08, "output_cost_per_token": 1.1e-06, "cache_read_input_token_cost": 7e-08, "litellm_provider": "openrouter", @@ -63919,8 +64346,8 @@ "supports_vision": false }, "openrouter/qwen/qwen3-30b-a3b-instruct-2507": { - "input_cost_per_token": 4.815e-08, - "output_cost_per_token": 1.9305e-07, + "input_cost_per_token": 9e-08, + "output_cost_per_token": 3e-07, "litellm_provider": "openrouter", "max_input_tokens": 262144, "max_output_tokens": 32000, @@ -64118,8 +64545,8 @@ "supports_vision": false }, "openrouter/qwen/qwen3-14b": { - "input_cost_per_token": 1.2e-07, - "output_cost_per_token": 2.4e-07, + "input_cost_per_token": 2.275e-07, + "output_cost_per_token": 9.1e-07, "litellm_provider": "openrouter", "max_input_tokens": 131072, "max_output_tokens": 16384, diff --git a/schema.prisma b/schema.prisma index 817df082d8c..7d521d54791 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1514,6 +1514,7 @@ model LiteLLM_AutoRouterSession { classifier_cost Float @default(0) classifier_cost_recorded_turns Int @default(0) tier_turns Json @default("{}") + baseline_models Json @default("{}") @@id([api_key, session_id, router_name]) @@index([last_turn_at], map: "idx_autorouter_session_last_turn") diff --git a/tests/e2e/ui/constants.ts b/tests/e2e/ui/constants.ts index 71774c95d24..158dea53b71 100644 --- a/tests/e2e/ui/constants.ts +++ b/tests/e2e/ui/constants.ts @@ -17,6 +17,8 @@ export const ARTIFACT_DIR = process.env.E2E_UI_ARTIFACT_DIR || "."; export const MOCK_PRESIDIO_URL = (process.env.E2E_MOCK_PRESIDIO_URL || "http://127.0.0.1:8091").replace(/\/+$/, ""); +export const PROPAGATION_TIMEOUT_MS = 90_000; + const storagePath = (name: string): string => path.join(ARTIFACT_DIR, name); // Storage state paths for each role diff --git a/tests/e2e/ui/tests/guardrails/presidioUserStory.spec.ts b/tests/e2e/ui/tests/guardrails/presidioUserStory.spec.ts index d4ed5308342..b4dc7c4e8da 100644 --- a/tests/e2e/ui/tests/guardrails/presidioUserStory.spec.ts +++ b/tests/e2e/ui/tests/guardrails/presidioUserStory.spec.ts @@ -1,8 +1,8 @@ -import { test, expect, type Page as PlaywrightPage } from "@playwright/test"; -import { ADMIN_STORAGE_PATH, MOCK_PRESIDIO_URL } from "../../constants"; +import { test as base, expect, type Page as PlaywrightPage } from "@playwright/test"; +import { ADMIN_STORAGE_PATH, MOCK_PRESIDIO_URL, PROPAGATION_TIMEOUT_MS } from "../../constants"; import { navigateToPage, dismissFeedbackPopup } from "../../helpers/navigation"; import { Page } from "../../fixtures/pages"; -import { CHAT_MODEL_A, masterKey, rootPath, waitForSpendLogByPrompt } from "../../helpers/traffic"; +import { CHAT_MODEL_A, masterKey, rootPath, uniqueSuffix, waitForSpendLog } from "../../helpers/traffic"; import { openPlayground, selectModel, sendButton, onlyVisible } from "../../helpers/playground"; const RAW_EMAIL = "jane.doe@example.com"; @@ -62,15 +62,63 @@ async function deleteGuardrail(page: PlaywrightPage, guardrailName: string): Pro await expect(page.getByText(`Guardrail "${guardrailName}" deleted successfully`)).toBeVisible({ timeout: 10_000 }); } +const test = base.extend<{ guardrailName: string }>({ + guardrailName: async ({ page }, use) => { + const name = `e2e-presidio-story-${uniqueSuffix()}`; + await createPresidioGuardrail(page, name); + try { + await use(name); + } finally { + await deleteGuardrail(page, name); + } + }, +}); + test.describe("Presidio PII guardrail, end to end from the dashboard", () => { + test.describe.configure({ timeout: 5 * 60_000 }); test.use({ storageState: ADMIN_STORAGE_PATH }); - test("masks PII sent from the Playground and shows the run in Logs", async ({ page, request }) => { - const guardrailName = `e2e-presidio-story-${Date.now()}`; + test("masks PII sent from the Playground and shows the run in Logs", async ({ page, request, guardrailName }) => { const marker = `case-ref-${Math.random().toString(36).slice(2, 10)}`; const prompt = `${marker}. Email me at ${RAW_EMAIL} or call ${RAW_PHONE}.`; - await createPresidioGuardrail(page, guardrailName); + await expect + .poll( + async () => { + const completion = await request.post(`${rootPath()}/v1/chat/completions`, { + headers: { Authorization: `Bearer ${masterKey()}` }, + data: { + model: CHAT_MODEL_A, + messages: [ + { role: "user", content: `readiness-${uniqueSuffix()}. Email ${RAW_EMAIL}; phone ${RAW_PHONE}.` }, + ], + guardrails: [guardrailName], + }, + }); + expect(completion.ok(), `guardrail readiness request failed: ${await completion.text()}`).toBe(true); + const { id }: { id: string } = await completion.json(); + expect(id).toBeTruthy(); + await waitForSpendLog(request, id); + const stored = await request.get(`${rootPath()}/spend/logs`, { + headers: { Authorization: `Bearer ${masterKey()}` }, + params: { request_id: id }, + }); + expect(stored.ok(), `guardrail readiness log read failed: ${stored.status()}`).toBe(true); + const body = await stored.text(); + return { + rawEmail: body.includes(RAW_EMAIL), + rawPhone: body.includes(RAW_PHONE), + maskedEmail: body.includes(""), + maskedPhone: body.includes(""), + }; + }, + { + message: `${guardrailName} never masked email and phone data in a completed request`, + timeout: PROPAGATION_TIMEOUT_MS, + intervals: [2_000], + }, + ) + .toEqual({ rawEmail: false, rawPhone: false, maskedEmail: true, maskedPhone: true }); await openPlayground(page); await selectModel(page, CHAT_MODEL_A); @@ -85,29 +133,31 @@ test.describe("Presidio PII guardrail, end to end from the dashboard", () => { const input = onlyVisible(page.getByPlaceholder("Type your message", { exact: false })); await expect(input).toBeVisible({ timeout: 15_000 }); - await expect - .poll( - async () => { - await input.fill(prompt); - await sendButton(page).click(); - const res = await request.get(`${rootPath()}/spend/logs`, { - headers: { Authorization: `Bearer ${masterKey()}` }, - }); - if (!res.ok()) return false; - const rows: { metadata?: { applied_guardrails?: string[] } }[] = await res.json(); - return (Array.isArray(rows) ? rows : []).some((row) => - (row.metadata?.applied_guardrails ?? []).includes(guardrailName), - ); - }, - { - message: `the playground never produced a request that ran ${guardrailName}`, - timeout: 90_000, - intervals: [5_000], - }, - ) - .toBe(true); - - const requestId = await waitForSpendLogByPrompt(request, marker); + await input.fill(prompt); + const responsePromise = page.waitForResponse( + (response) => + response.request().method() === "POST" && + new URL(response.url()).pathname.endsWith("/chat/completions") && + (response.request().postData()?.includes(marker) ?? false), + ); + await sendButton(page).click(); + const response = await responsePromise; + expect(response.ok(), `Playground completion failed: ${response.status()}`).toBe(true); + const responseBody = await response.text(); + const chunks: { id?: string; error?: unknown }[] = response.headers()["content-type"]?.includes("text/event-stream") + ? responseBody + .split(/\r?\n/) + .filter((line) => line.startsWith("data: ") && line.trim() !== "data: [DONE]") + .map((line) => JSON.parse(line.slice(6))) + : [JSON.parse(responseBody)]; + expect( + chunks.some((chunk) => chunk.error), + "the Playground stream returned an error", + ).toBe(false); + const requestIds = [...new Set(chunks.map((chunk) => chunk.id).filter((id): id is string => !!id))]; + expect(requestIds, "the Playground response identifies exactly one completion").toHaveLength(1); + const [requestId] = requestIds; + await waitForSpendLog(request, requestId); const stored = await request.get(`${rootPath()}/spend/logs?request_id=${requestId}`, { headers: { Authorization: `Bearer ${masterKey()}` }, @@ -138,7 +188,9 @@ test.describe("Presidio PII guardrail, end to end from the dashboard", () => { const drawer = page.getByRole("dialog").first(); await expect(onlyVisible(drawer.getByText("Guardrails & Policy Compliance"))).toBeVisible({ timeout: 20_000 }); - await expect(onlyVisible(drawer.getByText(`Pre-call guardrail: ${guardrailName}`))).toBeVisible({ timeout: 20_000 }); + await expect(onlyVisible(drawer.getByText(`Pre-call guardrail: ${guardrailName}`))).toBeVisible({ + timeout: 20_000, + }); const maskedPrompt = drawer.getByText(`${marker}. Email me at or call .`); await expect(onlyVisible(maskedPrompt)).toBeVisible({ timeout: 20_000 }); @@ -151,7 +203,5 @@ test.describe("Presidio PII guardrail, end to end from the dashboard", () => { await expect(drawer.getByText(RAW_EMAIL)).toHaveCount(0); await expect(drawer.getByText(RAW_PHONE)).toHaveCount(0); - - await deleteGuardrail(page, guardrailName); }); }); diff --git a/tests/e2e/ui/tests/modelsPage/modelHealthStatus.spec.ts b/tests/e2e/ui/tests/modelsPage/modelHealthStatus.spec.ts index 247cce1b85d..7bd1caa6756 100644 --- a/tests/e2e/ui/tests/modelsPage/modelHealthStatus.spec.ts +++ b/tests/e2e/ui/tests/modelsPage/modelHealthStatus.spec.ts @@ -4,7 +4,7 @@ import { type Locator, type Page as PlaywrightPage, } from "@playwright/test"; -import { ADMIN_STORAGE_PATH } from "../../constants"; +import { ADMIN_STORAGE_PATH, PROPAGATION_TIMEOUT_MS } from "../../constants"; import { Page } from "../../fixtures/pages"; import { navigateToPage } from "../../helpers/navigation"; import { readBack } from "../../helpers/roundTrip"; @@ -145,6 +145,23 @@ async function withDeployment( timeout: 60_000, }) .toBe(true); + await expect + .poll( + async () => { + const response = await page.request.get("/v1/models", { + headers: { Authorization: `Bearer ${masterKey()}` }, + }); + expect(response.ok(), `/v1/models failed: ${response.status()}`).toBe(true); + const body: { data: { id: string }[] } = await response.json(); + return body.data.some((model) => model.id === name); + }, + { + message: `deployment ${name} never appeared on the serving path`, + timeout: PROPAGATION_TIMEOUT_MS, + intervals: [2_000], + }, + ) + .toBe(true); await use(name); } finally { await deleteDeployment(page, id); @@ -161,6 +178,7 @@ const test = base.extend<{ reachableName: string; unreachableName: string }>({ }); test.describe("Model health status", () => { + test.describe.configure({ timeout: 8 * 60_000 }); test.use({ storageState: ADMIN_STORAGE_PATH }); test("Run Health Check reports a reachable deployment healthy and an unreachable one unhealthy", async ({ diff --git a/tests/e2e/ui/tests/proxy-admin/keyBlocking.spec.ts b/tests/e2e/ui/tests/proxy-admin/keyBlocking.spec.ts index 99a8065a797..788f4ca1284 100644 --- a/tests/e2e/ui/tests/proxy-admin/keyBlocking.spec.ts +++ b/tests/e2e/ui/tests/proxy-admin/keyBlocking.spec.ts @@ -1,5 +1,5 @@ import { test as base, expect } from "@playwright/test"; -import { ADMIN_STORAGE_PATH } from "../../constants"; +import { ADMIN_STORAGE_PATH, PROPAGATION_TIMEOUT_MS } from "../../constants"; import { Page } from "../../fixtures/pages"; import { dismissFeedbackPopup, navigateToPage, openKeyDetail } from "../../helpers/navigation"; import { @@ -35,6 +35,7 @@ test.describe("Proxy Admin - Key blocking", () => { test.use({ storageState: ADMIN_STORAGE_PATH }); test("blocking a key stops it serving and unblocking restores it", async ({ page, scopedKey }) => { + test.setTimeout(5 * 60_000); const { alias, token, apiKey } = scopedKey; await sendChatCompletion(page.request, { @@ -70,7 +71,7 @@ test.describe("Proxy Admin - Key blocking", () => { }), { message: "a blocked key was still served by /v1/chat/completions", - timeout: 30_000, + timeout: PROPAGATION_TIMEOUT_MS, }, ) .toMatchObject({ status: 401, body: expect.stringContaining("blocked") }); @@ -104,7 +105,7 @@ test.describe("Proxy Admin - Key blocking", () => { }), { message: "an unblocked key is still refused by /v1/chat/completions", - timeout: 30_000, + timeout: PROPAGATION_TIMEOUT_MS, }, ) .toMatchObject({ status: 200, body: expect.stringContaining(MOCK_RESPONSE_TEXT) }); diff --git a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py index 7e338dafb86..4a392042d63 100644 --- a/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py +++ b/tests/openai_endpoints_tests/test_e2e_openai_responses_api.py @@ -127,14 +127,14 @@ def test_bad_request_bad_param_error(): ) -def test_anthropic_with_responses_api(): - client = get_test_client() - response = client.responses.create( - model="anthropic/claude-sonnet-4-5-20250929", +def test_anthropic_with_responses_api() -> None: + client: Final = get_test_client() + response: Final = client.responses.create( + model="anthropic/claude-sonnet-5", input="just respond with the word 'ping'", - previous_response_id="hi", ) - print("anthropic response=", response) + assert response.status == "completed" + assert response.output_text.strip() def test_cancel_response(): diff --git a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py index 9ac6476a03c..e8caa241a53 100644 --- a/tests/proxy_behavior/spend/test_autorouter_session_rollup.py +++ b/tests/proxy_behavior/spend/test_autorouter_session_rollup.py @@ -42,6 +42,7 @@ async def _turn( saved: float = 0.02, classifier_cost: float = 0.0, tier: "str | None" = None, + baseline: "str | None" = None, ) -> None: touched: Final = 1 if (hit or ttl is not None or not covered) else 0 await db.execute_raw( @@ -61,6 +62,7 @@ async def _turn( ttl, touched, tier, + baseline, ) @@ -338,6 +340,27 @@ async def test_a_mid_session_router_type_change_keeps_foreign_tier_names_out_of_ assert row["turns"] == 3 +async def test_baseline_models_count_the_turns_priced_against_each_baseline(db): + key = f"k-{uuid.uuid4()}" + await _turn(db, key, "A", T0, baseline="opus") + await _turn(db, key, "B", T0 + timedelta(seconds=10), baseline="opus") + await _turn(db, key, "A", T0 + timedelta(seconds=20), baseline="sonnet") + + assert (await _row(db, key))["baseline_models"] == {"opus": 2, "sonnet": 1} + + +async def test_a_turn_priced_against_no_baseline_leaves_the_map_alone(db): + key = f"k-{uuid.uuid4()}" + await _turn(db, key, "A", T0, baseline=None) + assert (await _row(db, key))["baseline_models"] == {} + + await _turn(db, key, "A", T0 + timedelta(seconds=10), baseline="opus") + await _turn(db, key, "A", T0 + timedelta(seconds=20), baseline=None) + row = await _row(db, key) + assert row["baseline_models"] == {"opus": 1} + assert row["turns"] == 3 + + async def test_an_out_of_order_turn_still_counts_toward_its_tier(db): key = f"k-{uuid.uuid4()}" await _turn(db, key, "A", T0 + timedelta(seconds=60), tier="simple") 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 713330a8280..3283af26cf7 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -1895,3 +1895,15 @@ async def test_optional_discovery_preserves_cancellation(method: str) -> None: task.cancel() with pytest.raises(asyncio.CancelledError): await asyncio.wait_for(task, timeout=3) + + + +def test_client_import_before_proxy_credentials_succeeds_in_fresh_process(): + import subprocess + + result = subprocess.run( + [sys.executable, "-c", "import litellm.experimental_mcp_client.client; from litellm.proxy._experimental.mcp_server.mcp_server_manager import MCPServerManager; print(MCPServerManager.__name__)"], + capture_output=True, text=True, timeout=60, check=False, + ) + assert result.returncode == 0, result.stderr + assert result.stdout.strip() == "MCPServerManager" diff --git a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py b/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py index 0555447e34f..a2f81091893 100644 --- a/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py +++ b/tests/test_litellm/integrations/datadog/test_datadog_llm_obs.py @@ -541,24 +541,166 @@ def _span_json(logger_under_test: DataDogLLMObsLogger, payload: dict[str, Any]) return json.loads(safe_dumps(span)) +SECRET_TOOL_RESULT: Final = '{"city": "Paris", "temp_c": 18, "account_secret": "SECRET-7545"}' +TOOL_CONVERSATION: Final[list[dict[str, Any]]] = [ + {"role": "user", "content": "secret prompt"}, + {"role": "assistant", "content": None, "tool_calls": [ASSISTANT_TOOL_CALL]}, + {"role": "tool", "tool_call_id": "call_abc123", "content": SECRET_TOOL_RESULT}, +] + + +def _redacted_span_as_the_proxy_builds_it(payload: dict[str, Any]) -> dict[str, Any]: + logger_under_test = _redacting_logger(turn_off_message_logging=True) + return _span_json( + logger_under_test, logger_under_test.redact_standard_logging_payload_from_model_call_details(payload) + ) + + def test_redaction_keeps_the_conversation_shape_without_its_content() -> None: - """Roles and message count survive so the trace stays legible; contents and tool payloads do not.""" - result = _span_json( - _redacting_logger(turn_off_message_logging=True), + result = _redacted_span_as_the_proxy_builds_it( + build_payload( + messages=TOOL_CONVERSATION, + response_message={"role": "assistant", "content": "secret response", "tool_calls": [ASSISTANT_TOOL_CALL]}, + ) + ) + + redacted_call = { + "name": "get_weather", + "arguments": "redacted-by-litellm", + "tool_id": "call_abc123", + "type": "function", + } + assert result["meta"]["input"]["messages"] == [ + {"role": "user", "content": "redacted-by-litellm"}, + {"role": "assistant", "content": "redacted-by-litellm", "tool_calls": [redacted_call]}, + { + "role": "tool", + "content": "redacted-by-litellm", + "tool_results": [ + {"name": "get_weather", "result": "redacted-by-litellm", "tool_id": "call_abc123", "type": "function"} + ], + }, + ] + assert result["meta"]["output"]["messages"] == [ + {"role": "assistant", "content": "redacted-by-litellm", "tool_calls": [redacted_call]} + ] + serialized = safe_dumps(result) + assert "SECRET-7545" not in serialized + assert "Paris" not in serialized + assert "secret" not in serialized + + +def test_redaction_counts_tool_result_tokens_before_replacing_them() -> None: + payload = build_payload(messages=TOOL_CONVERSATION) + payload["standard_logging_object"]["model"] = "claude-sonnet-5" + + result = _redacted_span_as_the_proxy_builds_it(payload) + + expected_tokens = litellm.token_counter(model="claude-sonnet-5", text=SECRET_TOOL_RESULT) + assert expected_tokens > 0 + assert result["metrics"]["tool_output_tokens"] == float(expected_tokens) + assert result["metrics"]["input_tokens"] == 4447.0 + + +def test_tool_output_tokens_sum_every_result_in_the_request(logger: DataDogLLMObsLogger) -> None: + payload = build( + logger, + messages=[ + {"role": "tool", "tool_call_id": "call_1", "content": "one two three"}, + {"role": "tool", "tool_call_id": "call_2", "content": "four five six seven"}, + ], + ) + + assert payload["metrics"]["tool_output_tokens"] == float( + litellm.token_counter(text="one two three") + litellm.token_counter(text="four five six seven") + ) + + +def test_a_request_without_tool_results_reports_no_tool_output_tokens(logger: DataDogLLMObsLogger) -> None: + payload = build(logger, messages=[{"role": "user", "content": "hi"}]) + + assert "tool_output_tokens" not in payload["metrics"] + assert "tool_output_tokens" not in _redacted_span_as_the_proxy_builds_it(build_payload())["metrics"] + + +def test_a_tool_that_returned_nothing_still_counts_as_zero_tool_output_tokens(logger: DataDogLLMObsLogger) -> None: + payload = build(logger, messages=[{"role": "tool", "tool_call_id": "call_1", "content": ""}]) + + assert payload["metrics"]["tool_output_tokens"] == 0.0 + + +def test_redaction_keeps_anthropic_tool_blocks_as_structure_only() -> None: + result = _redacted_span_as_the_proxy_builds_it( build_payload( messages=[ - {"role": "user", "content": "secret prompt"}, - {"role": "assistant", "content": None, "tool_calls": [ASSISTANT_TOOL_CALL]}, - ], - response_message={"role": "assistant", "content": "secret response"}, - ), + { + "role": "assistant", + "content": [ + {"type": "tool_use", "id": "toolu_1", "name": "get_weather", "input": {"city": "Paris"}} + ], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": "toolu_1", "content": SECRET_TOOL_RESULT}], + }, + ] + ) ) assert result["meta"]["input"]["messages"] == [ - {"role": "user", "content": "redacted-by-litellm"}, - {"role": "assistant", "content": "redacted-by-litellm"}, + { + "role": "assistant", + "content": "redacted-by-litellm", + "tool_calls": [ + {"name": "get_weather", "arguments": "redacted-by-litellm", "tool_id": "toolu_1", "type": "tool_use"} + ], + }, + { + "role": "user", + "content": "redacted-by-litellm", + "tool_results": [ + {"name": "get_weather", "result": "redacted-by-litellm", "tool_id": "toolu_1", "type": "function"} + ], + }, ] - assert result["meta"]["output"]["messages"] == [{"role": "assistant", "content": "redacted-by-litellm"}] + assert "Paris" not in safe_dumps(result) + + +def test_redaction_blanks_tool_identifiers_that_are_not_strings() -> None: + result = _redacted_span_as_the_proxy_builds_it( + build_payload( + messages=[ + { + "role": "assistant", + "content": None, + "tool_calls": [ + {"id": {"leak": "SECRET-7545"}, "type": ["SECRET-7545"], "function": {"name": ["SECRET-7545"]}} + ], + } + ] + ) + ) + + assert result["meta"]["input"]["messages"][0]["tool_calls"] == [ + {"name": "", "arguments": "redacted-by-litellm", "tool_id": "", "type": ""} + ] + assert "SECRET-7545" not in safe_dumps(result) + + +def test_the_shared_hook_still_strips_what_redaction_governs_besides_messages() -> None: + payload = build_payload(messages=TOOL_CONVERSATION) + payload["standard_logging_object"]["classifier_input"] = {"system": "SECRET-7545"} + logger_under_test = _redacting_logger(turn_off_message_logging=True) + + with patch.object( # test-quality-ok: the hook reads this module global with no injection seam + litellm, "standard_logging_payload_excluded_fields", ["response"] + ): + redacted = logger_under_test.redact_standard_logging_payload_from_model_call_details(payload) + + assert "classifier_input" not in redacted["standard_logging_object"] + assert "response" not in redacted["standard_logging_object"] + assert redacted["standard_logging_object"]["messages"] == TOOL_CONVERSATION + assert payload["standard_logging_object"]["classifier_input"] == {"system": "SECRET-7545"} def test_redaction_drops_unrecognized_and_malformed_message_roles() -> None: diff --git a/tests/test_litellm/integrations/test_custom_guardrail.py b/tests/test_litellm/integrations/test_custom_guardrail.py index bdb7fad21b3..8cd719bfd55 100644 --- a/tests/test_litellm/integrations/test_custom_guardrail.py +++ b/tests/test_litellm/integrations/test_custom_guardrail.py @@ -2878,3 +2878,131 @@ async def test_native_lifecycle_guardrail_logging_only_scans_assembled_response( assert out_result is sentinel_result assert ("response", ["assembled stream text"]) in guardrail.calls assert out_kwargs["standard_logging_object"]["guardrail_information"] + + +class TestPreCallHookResponseIsNotLoggedVerbatim: + """Regression for LIT-6935: a pre_call hook returning the request payload leaked the prompt + into ``guardrail_response`` and from there onto OTEL guardrail spans.""" + + @staticmethod + def _logged_response(request_data: dict[str, object]) -> object: + metadata = request_data["litellm_metadata"] + assert isinstance(metadata, dict) + entries = metadata["standard_logging_guardrail_information"] + assert len(entries) == 1 + return entries[0]["guardrail_response"] + + @staticmethod + def _request() -> dict[str, object]: + return { + "model": "gpt-4.1-mini", + "input": "SECRET_PROMPT", + "messages": [{"role": "user", "content": "SECRET_PROMPT"}], + "litellm_metadata": {}, + } + + @pytest.mark.asyncio + async def test_pre_call_hook_returning_request_logs_allow(self): + class PassthroughGuardrail(CustomGuardrail): + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: object, + data: dict[str, object], + call_type: str, + ) -> dict[str, object]: + return data + + data = self._request() + await PassthroughGuardrail(guardrail_name="g").async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), cache=None, data=data, call_type="aresponses" + ) + + assert self._logged_response(data) == "allow" + + @pytest.mark.asyncio + async def test_pre_call_hook_returning_modified_copy_logs_mask(self): + class MaskingGuardrail(CustomGuardrail): + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: object, + data: dict[str, object], + call_type: str, + ) -> dict[str, object]: + return {**data, "input": "[MASKED]"} + + data = self._request() + await MaskingGuardrail(guardrail_name="g").async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), cache=None, data=data, call_type="aresponses" + ) + + assert self._logged_response(data) == "mask" + + @pytest.mark.asyncio + async def test_pre_call_hook_mutating_request_in_place_logs_mask(self): + class InPlaceMaskingGuardrail(CustomGuardrail): + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: object, + data: dict[str, object], + call_type: str, + ) -> dict[str, object]: + messages = data["messages"] + assert isinstance(messages, list) + messages[0]["content"] = "[MASKED]" + return data + + data = self._request() + await InPlaceMaskingGuardrail(guardrail_name="g").async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), cache=None, data=data, call_type="acompletion" + ) + + assert self._logged_response(data) == "mask" + + @pytest.mark.asyncio + async def test_pre_call_hook_returning_rejection_string_logs_that_string(self): + class RejectingGuardrail(CustomGuardrail): + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: object, + data: dict[str, object], + call_type: str, + ) -> str: + return "Blocked by policy" + + data = self._request() + result = await RejectingGuardrail(guardrail_name="g").async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), cache=None, data=data, call_type="acompletion" + ) + + assert result == "Blocked by policy" + assert self._logged_response(data) == "Blocked by policy" + + @pytest.mark.asyncio + async def test_pre_call_hook_removing_legacy_functions_in_place_logs_mask(self): + class FunctionStrippingGuardrail(CustomGuardrail): + @log_guardrail_information + async def async_pre_call_hook( + self, + user_api_key_dict: UserAPIKeyAuth, + cache: object, + data: dict[str, object], + call_type: str, + ) -> dict[str, object]: + data["functions"] = [] + data["function_call"] = "none" + return data + + data = {**self._request(), "functions": [{"name": "delete_db"}], "function_call": "auto"} + await FunctionStrippingGuardrail(guardrail_name="g").async_pre_call_hook( + user_api_key_dict=UserAPIKeyAuth(), cache=None, data=data, call_type="acompletion" + ) + + assert self._logged_response(data) == "mask" diff --git a/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py b/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py index ae5cffd8ab0..f1f9f7c3f3f 100644 --- a/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py +++ b/tests/test_litellm/integrations/vector_store_integrations/test_vector_store_pre_call_hook.py @@ -11,7 +11,16 @@ from litellm.integrations.vector_store_integrations.vector_store_pre_call_hook i ProxyServerRuntime, VectorStorePreCallHook, ) -from litellm.types.llms.openai import AllMessageValues +from litellm.types.llms.openai import AllMessageValues, ResponsesAPIResponse +from litellm.types.utils import ( + CallTypes, + Choices, + Delta, + Message, + ModelResponse, + ModelResponseStream, + StreamingChoices, +) from litellm.types.vector_stores import ( VectorStoreResultContent, VectorStoreSearchResponse, @@ -36,6 +45,18 @@ def _search_response(text: str) -> VectorStoreSearchResponse: ) +def _first_message(response: ModelResponse) -> Message: + choice = response.choices[0] + assert isinstance(choice, Choices) + return choice.message + + +@dataclass(frozen=True) +class ExplodingRegistry: + async def pop_vector_stores_to_run_with_db_fallback(self, **kwargs: object) -> list[LiteLLM_ManagedVectorStore]: + raise RuntimeError("the registry blew up") + + @dataclass class RecordingRouter: failing_vector_store_ids: frozenset[str] = frozenset() @@ -285,3 +306,264 @@ def test_the_default_runtime_follows_the_proxy_globals(monkeypatch: pytest.Monke assert runtime.llm_router() is router assert runtime.prisma_client() is prisma + + +@pytest.mark.asyncio +async def test_a_failing_vector_store_is_reported_back_to_the_caller( + registry_with: RegisterStores, +) -> None: + """Regression (LIT-6809): a silently dropped store left the caller with an un-augmented answer and no signal.""" + registry_with("vs-broken", "vs-healthy") + logging_obj = FakeLoggingObj({}) + + await _run_hook( + VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"}))) + ), + ["vs-broken", "vs-healthy"], + logging_obj, + ) + + response = ModelResponse(choices=[Choices(message=Message(content="an answer"))]) + await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)).async_post_call_success_deployment_hook( + request_data={"litellm_logging_obj": logging_obj}, + response=response, + call_type=CallTypes.acompletion, + ) + + provider_specific_fields = _first_message(response).provider_specific_fields or {} + assert provider_specific_fields["vector_store_search_failures"] == ( + { + "vector_store_id": "vs-broken", + "custom_llm_provider": "bedrock", + "error": "litellm.BadRequestError: no healthy deployments for vs-broken", + }, + ) + assert len(provider_specific_fields["search_results"]) == 1 + + +@pytest.mark.asyncio +async def test_a_healthy_vector_store_alone_reports_no_failures(registry_with: RegisterStores) -> None: + registry_with("vs-healthy") + logging_obj = FakeLoggingObj({}) + + await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=RecordingRouter())), + ["vs-healthy"], + logging_obj, + ) + + response = ModelResponse(choices=[Choices(message=Message(content="an answer"))]) + await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)).async_post_call_success_deployment_hook( + request_data={"litellm_logging_obj": logging_obj}, + response=response, + call_type=CallTypes.acompletion, + ) + + assert "vector_store_search_failures" not in (_first_message(response).provider_specific_fields or {}) + + +@pytest.mark.asyncio +async def test_a_failing_vector_store_is_reported_on_the_responses_api_response( + registry_with: RegisterStores, +) -> None: + """Regression (LIT-6809): /v1/responses answered 200 with no sign the knowledge base was missing.""" + registry_with("vs-broken") + logging_obj = FakeLoggingObj({}) + + await _run_hook( + VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"}))) + ), + ["vs-broken"], + logging_obj, + ) + + response = ResponsesAPIResponse(id="resp-lit6809", created_at=0, output=[]) + await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)).async_post_call_success_deployment_hook( + request_data={"litellm_logging_obj": logging_obj}, + response=response, + call_type=CallTypes.aresponses, + ) + + assert response.model_dump()["vector_store_search_failures"] == [ + { + "vector_store_id": "vs-broken", + "custom_llm_provider": "bedrock", + "error": "litellm.BadRequestError: no healthy deployments for vs-broken", + } + ] + + +@pytest.mark.asyncio +async def test_a_healthy_vector_store_leaves_the_responses_api_response_alone(registry_with: RegisterStores) -> None: + registry_with("vs-healthy") + logging_obj = FakeLoggingObj({}) + + await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=RecordingRouter())), + ["vs-healthy"], + logging_obj, + ) + + response = ResponsesAPIResponse(id="resp-lit6809", created_at=0, output=[]) + await VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)).async_post_call_success_deployment_hook( + request_data={"litellm_logging_obj": logging_obj}, + response=response, + call_type=CallTypes.aresponses, + ) + + assert "vector_store_search_failures" not in response.model_dump() + + +@pytest.mark.asyncio +async def test_a_failing_vector_store_is_reported_on_the_streaming_chunk(registry_with: RegisterStores) -> None: + registry_with("vs-broken") + logging_obj = FakeLoggingObj({}) + + await _run_hook( + VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"}))) + ), + ["vs-broken"], + logging_obj, + ) + + chunk = ModelResponseStream(choices=[StreamingChoices(delta=Delta(content="an answer"))]) + await VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=None) + ).async_post_call_streaming_deployment_hook( + request_data=logging_obj.model_call_details, + response_chunk=chunk, + call_type=CallTypes.acompletion, + ) + + assert (chunk.choices[0].delta.provider_specific_fields or {})["vector_store_search_failures"] == ( + { + "vector_store_id": "vs-broken", + "custom_llm_provider": "bedrock", + "error": "litellm.BadRequestError: no healthy deployments for vs-broken", + }, + ) + + +@pytest.mark.asyncio +async def test_error_mode_fails_the_request_instead_of_answering_without_the_knowledge_base( + registry_with: RegisterStores, + monkeypatch: pytest.MonkeyPatch, +) -> None: + """Regression (LIT-6809): opting in must turn an ungrounded answer into a 400 the caller can act on.""" + registry_with("vs-broken", "vs-healthy") + monkeypatch.setattr(litellm, "vector_store_search_failure_mode", "error") + + with pytest.raises(litellm.VectorStoreSearchError) as raised: + await _run_hook( + VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime( + router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"})) + ) + ), + ["vs-broken", "vs-healthy"], + FakeLoggingObj({}), + ) + + assert raised.value.status_code == 400 + assert raised.value.failures == ( + { + "vector_store_id": "vs-broken", + "custom_llm_provider": "bedrock", + "error": "litellm.BadRequestError: no healthy deployments for vs-broken", + }, + ) + assert "vs-broken: litellm.BadRequestError: no healthy deployments for vs-broken" in raised.value.message + + +@pytest.mark.asyncio +async def test_a_misspelled_failure_mode_annotates_instead_of_erroring_the_request( + registry_with: RegisterStores, + monkeypatch: pytest.MonkeyPatch, + warnings: list[logging.LogRecord], +) -> None: + """Regression (LIT-6809): litellm_settings takes any value, so a typo must not become a 500.""" + registry_with("vs-broken") + monkeypatch.setattr(litellm, "vector_store_search_failure_mode", "erorr") + + logging_obj = FakeLoggingObj({}) + _, messages, _ = await _run_hook( + VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime(router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"}))) + ), + ["vs-broken"], + logging_obj, + ) + + assert messages[0]["content"] == "what is litellm?" + assert logging_obj.model_call_details["vector_store_search_failures"] == ( + { + "vector_store_id": "vs-broken", + "custom_llm_provider": "bedrock", + "error": "litellm.BadRequestError: no healthy deployments for vs-broken", + }, + ) + assert any("erorr" in record.getMessage() for record in warnings) + + +@pytest.mark.asyncio +async def test_error_mode_leaves_a_fully_healthy_request_alone( + registry_with: RegisterStores, + monkeypatch: pytest.MonkeyPatch, +) -> None: + registry_with("vs-healthy") + monkeypatch.setattr(litellm, "vector_store_search_failure_mode", "error") + + _, messages, _ = await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=RecordingRouter())), + ["vs-healthy"], + FakeLoggingObj({}), + ) + + assert messages[0]["content"] == "Context:\n\ncontext from vs-healthy\n\n" + + +@pytest.mark.asyncio +async def test_error_mode_does_not_swallow_the_raise_in_the_hooks_own_catch_all( + registry_with: RegisterStores, + monkeypatch: pytest.MonkeyPatch, + warnings: list[logging.LogRecord], +) -> None: + """Regression (LIT-6809): the catch-all around the hook must not turn the opted-in failure back into a 200.""" + registry_with("vs-broken") + monkeypatch.setattr(litellm, "vector_store_search_failure_mode", "error") + + with pytest.raises(litellm.VectorStoreSearchError): + await _run_hook( + VectorStorePreCallHook( + proxy_runtime=FakeProxyRuntime( + router=RecordingRouter(failing_vector_store_ids=frozenset({"vs-broken"})) + ) + ), + ["vs-broken"], + FakeLoggingObj({}), + ) + + assert [record.levelname for record in warnings] == ["WARNING"] + + +@pytest.mark.asyncio +async def test_a_crash_outside_the_search_names_the_requested_vector_stores( + monkeypatch: pytest.MonkeyPatch, + warnings: list[logging.LogRecord], +) -> None: + """Regression (LIT-6809): the catch-all logged no store id, so an operator could not tell which store broke.""" + monkeypatch.setattr(litellm, "vector_store_registry", ExplodingRegistry()) + + _, messages, _ = await _run_hook( + VectorStorePreCallHook(proxy_runtime=FakeProxyRuntime(router=None)), + ["vs-one", "vs-two"], + FakeLoggingObj({}), + ) + + assert messages == [{"role": "user", "content": "what is litellm?"}] + assert [record.getMessage() for record in warnings] == [ + "Error in VectorStorePreCallHook for vector_store_ids=('vs-one', 'vs-two'): the registry blew up" + ] diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index fbb9d178390..cbe6fe198c9 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -1715,7 +1715,7 @@ def test_generic_cost_per_token_gpt55(_local_model_cost_map): def test_generic_cost_per_token_gpt55_pro(_local_model_cost_map): - """gpt-5.5-pro: responses-only model — $30/1M input, $180/1M output, $3/1M cached input.""" + """gpt-5.5-pro: responses-only model, $30/1M input, $180/1M output, no cached input rate published.""" model = "gpt-5.5-pro" custom_llm_provider = "openai" @@ -1724,7 +1724,7 @@ def test_generic_cost_per_token_gpt55_pro(_local_model_cost_map): # Sanity-check the map values match OpenAI's published pricing. assert model_cost_map["input_cost_per_token"] == 3e-5 assert model_cost_map["output_cost_per_token"] == 1.8e-4 - assert model_cost_map["cache_read_input_token_cost"] == 3e-6 + assert "cache_read_input_token_cost" not in model_cost_map assert model_cost_map["litellm_provider"] == "openai" # gpt-5.5-pro is a responses-only model (no /v1/chat/completions endpoint). assert model_cost_map["mode"] == "responses" diff --git a/tests/test_litellm/litellm_core_utils/test_private_json.py b/tests/test_litellm/litellm_core_utils/test_private_json.py index cedff61959f..3c9f49607f9 100644 --- a/tests/test_litellm/litellm_core_utils/test_private_json.py +++ b/tests/test_litellm/litellm_core_utils/test_private_json.py @@ -4,7 +4,11 @@ import stat import pytest -from litellm.litellm_core_utils.private_json import overwrite_private_json, write_private_json +from litellm.litellm_core_utils.private_json import ( + overwrite_private_json, + write_private_bytes, + write_private_json, +) class TestOverwritePrivateJson: @@ -35,3 +39,32 @@ class TestOverwritePrivateJson: overwrite_private_json(str(path), {"user_id": "u-1"}) assert stat.S_IMODE(path.stat().st_mode) == 0o600 + + +class TestWritePrivateBytes: + def test_replaces_the_file_in_one_step_so_a_reader_holding_the_old_one_keeps_it_whole(self, tmp_path): + path = tmp_path / "script.py" + write_private_bytes(str(path), b"print('one')\n" * 200) + before = path.stat().st_ino + + with path.open("rb") as reader: + write_private_bytes(str(path), b"print('two')\n") + assert reader.read() == b"print('one')\n" * 200 + + assert path.read_bytes() == b"print('two')\n" + assert path.stat().st_ino != before + assert [child.name for child in tmp_path.iterdir()] == ["script.py"] + + @pytest.mark.skipif(os.geteuid() == 0, reason="root ignores file permissions") + def test_lands_owner_only_and_a_refused_stage_leaves_the_previous_file_untouched(self, tmp_path): + path = tmp_path / "script.py" + write_private_bytes(str(path), b"first") + assert stat.S_IMODE(path.stat().st_mode) == 0o600 + + tmp_path.chmod(0o500) + try: + with pytest.raises(PermissionError): + write_private_bytes(str(path), b"second") + finally: + tmp_path.chmod(0o700) + assert path.read_bytes() == b"first" diff --git a/tests/test_litellm/litellm_core_utils/test_redact_messages.py b/tests/test_litellm/litellm_core_utils/test_redact_messages.py index 76092e96307..584a3ac471c 100644 --- a/tests/test_litellm/litellm_core_utils/test_redact_messages.py +++ b/tests/test_litellm/litellm_core_utils/test_redact_messages.py @@ -926,3 +926,38 @@ def test_classifier_callback_redaction_preserves_exclusions(monkeypatch: pytest. assert failure_payload["response"]["choices"][0]["message"]["content"] == "redacted-by-litellm" assert payload["classifier_input"] == {"system": "private rubric"} assert payload["response"]["choices"][0]["message"]["content"] == "private answer" + + +class _SelfRedactingLogger(CustomLogger): + def redacts_messages_itself(self) -> bool: + return True + + +@pytest.mark.parametrize("logger", [CustomLogger(), _SelfRedactingLogger()], ids=["default", "redacts_itself"]) +def test_field_exclusion_alone_leaves_messages_and_responses_intact(monkeypatch: pytest.MonkeyPatch, logger: CustomLogger) -> None: + monkeypatch.setattr(litellm, "standard_logging_payload_excluded_fields", ["model"]) + payload: Final = { + "messages": [{"role": "user", "content": "private prompt"}], + "response": {"choices": [{"message": {"content": "private answer"}}]}, + "model": "classifier", + } + stored: Final = logger.redact_standard_logging_payload_from_model_call_details({"standard_logging_object": payload})[ + "standard_logging_object" + ] + assert stored == {"messages": payload["messages"], "response": payload["response"]} + + +def test_a_callback_that_redacts_itself_keeps_its_messages_but_not_the_classifier_audit() -> None: + payload: Final = { + "classifier_input": {"system": "private rubric"}, + "messages": [{"role": "user", "content": "private prompt"}], + "response": {"choices": [{"message": {"content": "private answer"}}]}, + } + logger: Final = _SelfRedactingLogger() + logger.turn_off_message_logging = True + stored: Final = logger.redact_standard_logging_payload_from_model_call_details({"standard_logging_object": payload})[ + "standard_logging_object" + ] + assert "classifier_input" not in stored + assert stored["messages"] == payload["messages"] + assert stored["response"] == payload["response"] diff --git a/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py b/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py index c8007acf70f..f00698a6624 100644 --- a/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py +++ b/tests/test_litellm/llms/azure_ai/passthrough/test_azure_ai_passthrough_transformation.py @@ -608,9 +608,11 @@ async def test_streaming_responses_relay_flush_reaches_the_success_callbacks_wit ) stream = "event: response.completed\ndata: " + json.dumps(RESPONSES_COMPLETED_EVENT) + "\n\n" - await logging_obj.async_flush_passthrough_collected_chunks( - raw_bytes=[stream.encode()], provider_config=AzureAIPassthroughConfig() + collector = AzureAIPassthroughConfig().create_stream_collector( + model="gpt-5.4-mini", custom_llm_provider="azure_ai", endpoint="gpt/openai/responses" ) + collector.add(stream.encode()) + await logging_obj.async_flush_passthrough_collected_chunks(collector=collector) info = litellm.get_model_info("azure_ai/gpt-5.4-mini") assert probe.logged_call_type == "allm_passthrough_route" diff --git a/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py b/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py index b005d77ac8b..f2a9af11af7 100644 --- a/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py +++ b/tests/test_litellm/llms/bedrock/passthrough/test_bedrock_passthrough_transformation.py @@ -1,7 +1,19 @@ +import base64 +import json +import struct +import tracemalloc +from binascii import crc32 +from datetime import datetime from unittest.mock import patch - +from litellm.litellm_core_utils.litellm_logging import Logging +from litellm.llms.base_llm.passthrough.transformation import PassthroughStreamCollector from litellm.llms.bedrock.passthrough.transformation import BedrockPassthroughConfig +from litellm.types.utils import ModelResponse + +CONVERSE_MODEL = "anthropic.claude-sonnet-4-5-20250929-v1:0" +CONVERSE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/converse-stream" +INVOKE_STREAM_ENDPOINT = f"/model/{CONVERSE_MODEL}/invoke-with-response-stream" def test_bedrock_passthrough_get_complete_url_default_endpoint(): @@ -500,3 +512,186 @@ def test_bedrock_passthrough_model_id_without_arn(): f"https://bedrock-runtime.us-east-1.amazonaws.com/model/{model_id}/converse" ) assert url_str == expected_url + + +def _event_frame(event_type: str, payload: dict) -> bytes: + def header(name: str, value: str) -> bytes: + name_b, value_b = name.encode(), value.encode() + return struct.pack("!B", len(name_b)) + name_b + struct.pack("!B", 7) + struct.pack("!H", len(value_b)) + value_b + + payload_b = json.dumps(payload, separators=(",", ":")).encode() + headers_b = ( + header(":event-type", event_type) + + header(":content-type", "application/json") + + header(":message-type", "event") + ) + prelude = struct.pack("!II", 12 + len(headers_b) + len(payload_b) + 4, len(headers_b)) + prelude_crc = crc32(prelude) & 0xFFFFFFFF + message = struct.pack("!I", prelude_crc) + headers_b + payload_b + return prelude + message + struct.pack("!I", crc32(message, prelude_crc) & 0xFFFFFFFF) + + +def _text_block(index: int, texts: list[str]) -> bytes: + return ( + _event_frame("contentBlockStart", {"contentBlockIndex": index, "start": {}}) + + b"".join( + _event_frame("contentBlockDelta", {"contentBlockIndex": index, "delta": {"text": text}}) for text in texts + ) + + _event_frame("contentBlockStop", {"contentBlockIndex": index}) + ) + + +def _stream_tail(stop_reason: str, output_tokens: int) -> bytes: + return _event_frame("messageStop", {"stopReason": stop_reason}) + _event_frame( + "metadata", + { + "metrics": {"latencyMs": 1234}, + "usage": {"inputTokens": 25, "outputTokens": output_tokens, "totalTokens": 25 + output_tokens}, + }, + ) + + +def _invoke_chunk(payload: dict) -> bytes: + return _event_frame("chunk", {"bytes": base64.b64encode(json.dumps(payload).encode()).decode()}) + + +def _stream_logging_obj(endpoint: str) -> Logging: + logging_obj = Logging( + model=CONVERSE_MODEL, + messages=[], + stream=True, + call_type="pass_through_endpoint", + start_time=datetime.now(), + litellm_call_id="call-1", + function_id="fn-1", + ) + logging_obj.model_call_details["custom_llm_provider"] = "bedrock" + logging_obj.model_call_details["endpoint"] = endpoint + return logging_obj + + +def _converse_stream_logging_obj() -> Logging: + return _stream_logging_obj(CONVERSE_STREAM_ENDPOINT) + + +def _stream_collector(endpoint: str) -> PassthroughStreamCollector: + return BedrockPassthroughConfig().create_stream_collector( + model=CONVERSE_MODEL, custom_llm_provider="bedrock", endpoint=endpoint + ) + + +def _converse_stream_collector() -> PassthroughStreamCollector: + return _stream_collector(CONVERSE_STREAM_ENDPOINT) + + +def _feed(collector: PassthroughStreamCollector, stream: bytes, chunk_size: int = 16384) -> None: + for offset in range(0, len(stream), chunk_size): + collector.add(stream[offset : offset + chunk_size]) + + +def test_converse_stream_collector_keeps_usage_without_retaining_the_stream(): + texts = [f"tok{i} " for i in range(4000)] + stream = _event_frame("messageStart", {"role": "assistant"}) + _text_block(0, texts) + _stream_tail("end_turn", 4000) + _feed(_converse_stream_collector(), stream) + + tracemalloc.start() + try: + base = tracemalloc.get_traced_memory()[0] + collector = _converse_stream_collector() + _feed(collector, stream) + retained = tracemalloc.get_traced_memory()[0] - base + finally: + tracemalloc.stop() + + assert retained < len(stream) // 4 + + response = collector.build_logged_response(_converse_stream_logging_obj()) + assert isinstance(response, ModelResponse) + assert response.choices[0].message.content == "".join(texts) + assert response.choices[0].finish_reason == "stop" + assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 4000) + + +def test_converse_stream_collector_keeps_tool_calls_between_text_runs(): + stream = ( + _event_frame("messageStart", {"role": "assistant"}) + + _text_block(0, ["Let me ", "check."]) + + _event_frame( + "contentBlockStart", + {"contentBlockIndex": 1, "start": {"toolUse": {"toolUseId": "tool-1", "name": "get_weather"}}}, + ) + + _event_frame("contentBlockDelta", {"contentBlockIndex": 1, "delta": {"toolUse": {"input": '{"city": '}}}) + + _event_frame("contentBlockDelta", {"contentBlockIndex": 1, "delta": {"toolUse": {"input": '"Paris"}'}}}) + + _event_frame("contentBlockStop", {"contentBlockIndex": 1}) + + _text_block(2, ["Done", "."]) + + _stream_tail("tool_use", 12) + ) + collector = _converse_stream_collector() + _feed(collector, stream, chunk_size=7) + + response = collector.build_logged_response(_converse_stream_logging_obj()) + assert isinstance(response, ModelResponse) + message = response.choices[0].message + assert message.content == "Let me check.Done." + assert [(call.function.name, call.function.arguments) for call in message.tool_calls] == [ + ("get_weather", '{"city": "Paris"}') + ] + assert response.choices[0].finish_reason == "tool_calls" + assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 12) + + +def test_invoke_stream_collector_keeps_usage_without_retaining_the_stream(): + texts = [f"tok{i} " for i in range(4000)] + stream = ( + _invoke_chunk( + { + "type": "message_start", + "message": { + "id": "msg-1", + "type": "message", + "role": "assistant", + "model": CONVERSE_MODEL, + "content": [], + "stop_reason": None, + "usage": {"input_tokens": 25, "output_tokens": 1}, + }, + } + ) + + _invoke_chunk({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}) + + b"".join( + _invoke_chunk({"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}) + for text in texts + ) + + _invoke_chunk({"type": "content_block_stop", "index": 0}) + + _invoke_chunk( + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 4000}} + ) + + _invoke_chunk({"type": "message_stop"}) + ) + _feed(_stream_collector(INVOKE_STREAM_ENDPOINT), stream) + + tracemalloc.start() + try: + base = tracemalloc.get_traced_memory()[0] + collector = _stream_collector(INVOKE_STREAM_ENDPOINT) + _feed(collector, stream) + retained = tracemalloc.get_traced_memory()[0] - base + finally: + tracemalloc.stop() + + assert retained < len(stream) // 4 + + response = collector.build_logged_response(_stream_logging_obj(INVOKE_STREAM_ENDPOINT)) + assert isinstance(response, ModelResponse) + assert response.choices[0].message.content == "".join(texts) + assert response.choices[0].finish_reason == "stop" + assert (response.usage.prompt_tokens, response.usage.completion_tokens) == (25, 4000) + + +def test_stream_collector_logs_nothing_for_an_unrecognized_endpoint(): + collector = BedrockPassthroughConfig().create_stream_collector( + model=CONVERSE_MODEL, custom_llm_provider="bedrock", endpoint=f"/model/{CONVERSE_MODEL}/rerank" + ) + collector.add(_event_frame("messageStart", {"role": "assistant"})) + + assert collector.build_logged_response(_converse_stream_logging_obj()) is None diff --git a/tests/test_litellm/llms/cerebras/test_cerebras_chat_transformation.py b/tests/test_litellm/llms/cerebras/test_cerebras_chat_transformation.py index 09718b1e6e0..a47180e9511 100644 --- a/tests/test_litellm/llms/cerebras/test_cerebras_chat_transformation.py +++ b/tests/test_litellm/llms/cerebras/test_cerebras_chat_transformation.py @@ -1,3 +1,6 @@ +import pytest + +import litellm from litellm.llms.cerebras.chat import CerebrasConfig @@ -59,3 +62,23 @@ def test_map_openai_params_preserves_max_retries_zero_falsy() -> None: assert "max_retries" in result and result["max_retries"] == 0, ( f"max_retries=0 (falsy) must not be silently omitted; got: {result!r}" ) + + +def test_qwen_3_8_27b_cost_and_tokens(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + model = "cerebras/qwen-3.8-27b" + prompt_cost, completion_cost = litellm.cost_per_token( + model=model, + prompt_tokens=1000, + completion_tokens=1000, + ) + assert abs(prompt_cost - 0.00099) < 1e-9 + assert abs(completion_cost - 0.00149) < 1e-9 + + model_info = litellm.get_model_info(model) + assert model_info["max_input_tokens"] == 65536 + assert model_info["max_output_tokens"] == 32768 + assert model_info["supports_vision"] is True + assert model_info["supports_reasoning"] is True + assert model_info["supports_parallel_function_calling"] is True diff --git a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py index cea2d439198..c39779972c0 100644 --- a/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py +++ b/tests/test_litellm/llms/custom_httpx/test_llm_http_handler.py @@ -3,6 +3,7 @@ import json import logging import threading import time +from typing import Final from unittest.mock import AsyncMock, Mock, patch import httpx @@ -20,7 +21,9 @@ from litellm.llms.base_llm.audio_transcription.transformation import ( BaseAudioTranscriptionConfig, ) from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException +from litellm.llms.base_llm.search.transformation import BaseSearchConfig, SearchResponse from litellm.llms.bedrock.base_aws_llm import SignsRequestsWithAWS +from litellm.llms.brave.search.transformation import BraveSearchConfig from litellm.llms.base_llm.image_edit.transformation import BaseImageEditConfig from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler from litellm.llms.custom_httpx.llm_http_handler import ( @@ -36,6 +39,7 @@ from litellm.llms.bedrock.messages.invoke_transformations.anthropic_claude3_tran ) from litellm.llms.mistral.ocr.transformation import MistralOCRConfig from litellm.llms.openai.videos.transformation import OpenAIVideoConfig +from litellm.llms.tinyfish.search.transformation import TinyfishSearchConfig from litellm.types.llms.openai import ResponsesAPIResponse from litellm.types.router import GenericLiteLLMParams from litellm.types.utils import ImageObject, ImageResponse, ModelResponse, TranscriptionResponse @@ -44,6 +48,95 @@ from tests.test_litellm.llms.bedrock.event_loop_probe import EventLoopProbe _ACTIVE_KEY = "_code_interpreter_interception_active" _SANDBOX_KEY = "_code_interpreter_interception_sandbox_key" + +async def _get_search_with_client( + client: HTTPHandler | AsyncHTTPHandler, provider_config: BaseSearchConfig | None = None +) -> SearchResponse: + result: Final = BaseLLMHTTPHandler().search( + query="test", + optional_params={}, + timeout=5, + logging_obj=Mock(), + api_key="test-key", + api_base="https://search.example.test/", + custom_llm_provider="tinyfish" if isinstance(provider_config, TinyfishSearchConfig) else "brave", + client=client, + asearch=isinstance(client, AsyncHTTPHandler), + provider_config=provider_config or BraveSearchConfig(), + ) + return await result if asyncio.iscoroutine(result) else result + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", (False, True)) +@pytest.mark.parametrize("status_code", (400, 401, 403, 422, 429, 500)) +async def test_get_search_raises_provider_http_errors(is_async: bool, status_code: int) -> None: + upstream_response: Final = httpx.Response( + status_code, json={"error": "rejected request"}, headers={"retry-after": "7"} + ) + transport: Final = httpx.MockTransport(lambda request: upstream_response) + async with httpx.AsyncClient(transport=transport) as async_client: + with httpx.Client(transport=transport) as sync_client: + client: Final = AsyncHTTPHandler() if is_async else HTTPHandler(client=sync_client) + if isinstance(client, AsyncHTTPHandler): + await client.close() + client.client = async_client + with pytest.raises(BaseLLMException) as error: + await _get_search_with_client(client) + assert error.value.status_code == status_code + assert "rejected request" in error.value.message + assert error.value.headers is not None + assert error.value.headers["retry-after"] == "7" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", (False, True)) +@pytest.mark.parametrize("has_results", (False, True)) +async def test_get_search_preserves_successful_results(is_async: bool, has_results: bool) -> None: + results: Final = ( + [{"title": "Example", "url": "https://example.com", "description": "Example snippet"}] if has_results else [] + ) + transport: Final = httpx.MockTransport(lambda request: httpx.Response(200, json={"web": {"results": results}})) + async with httpx.AsyncClient(transport=transport) as async_client: + with httpx.Client(transport=transport) as sync_client: + client: Final = AsyncHTTPHandler() if is_async else HTTPHandler(client=sync_client) + if isinstance(client, AsyncHTTPHandler): + await client.close() + client.client = async_client + response: Final = await _get_search_with_client(client) + assert response.object == "search" + assert len(response.results) == int(has_results) + if has_results: + assert response.results[0].title == "Example" + assert response.results[0].url == "https://example.com" + assert response.results[0].snippet == "Example snippet" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("is_async", (False, True)) +async def test_get_search_preserves_tinyfish_http_error_formatting(is_async: bool) -> None: + upstream_response: Final = httpx.Response( + 429, + json={"error": {"code": "RATE_LIMIT_EXCEEDED", "message": "rate limit exceeded"}}, + headers={"retry-after": "7"}, + ) + transport: Final = httpx.MockTransport(lambda request: upstream_response) + async with httpx.AsyncClient(transport=transport) as async_client: + with httpx.Client(transport=transport) as sync_client: + client: Final = AsyncHTTPHandler() if is_async else HTTPHandler(client=sync_client) + if isinstance(client, AsyncHTTPHandler): + await client.close() + client.client = async_client + with pytest.raises(BaseLLMException) as error: + await _get_search_with_client(client, TinyfishSearchConfig()) + assert error.value.status_code == 429 + assert error.value.message == ( + "TinyFish Search: rate limit exceeded. See https://docs.tinyfish.ai/search-api for details." + ) + assert error.value.headers is not None + assert error.value.headers["retry-after"] == "7" + + OCR_RESPONSE = { "pages": [{"index": 0, "markdown": "OCR output", "images": []}], "model": "mistral-ocr-latest", diff --git a/tests/test_litellm/llms/inception/test_inception_chat_transformation.py b/tests/test_litellm/llms/inception/test_inception_chat_transformation.py index 4c0f5969249..04813143fae 100644 --- a/tests/test_litellm/llms/inception/test_inception_chat_transformation.py +++ b/tests/test_litellm/llms/inception/test_inception_chat_transformation.py @@ -7,6 +7,7 @@ import os from unittest import mock import httpx +import pytest import litellm from litellm.llms.inception.chat.transformation import InceptionChatConfig @@ -238,6 +239,7 @@ def test_inception_model_list_populated(monkeypatch): litellm.add_known_models() assert "inception/mercury-2" in litellm.inception_models + assert "inception/mercury-2.5" in litellm.inception_models for model in litellm.inception_models: assert model.startswith("inception/") @@ -304,3 +306,24 @@ def test_inception_completion_targets_inception_endpoint(): assert captured["body"]["model"] == "mercury-2" assert captured["body"]["tool_choice"] == "auto" assert response.choices[0].message.content == "hi" + + +def test_inception_mercury_2_5_cost_and_tokens(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True") + monkeypatch.setattr(litellm, "model_cost", litellm.get_model_cost_map(url="")) + model = "inception/mercury-2.5" + prompt_cost, completion_cost = litellm.cost_per_token( + model=model, + prompt_tokens=1000, + completion_tokens=500, + ) + assert abs(prompt_cost - 0.0002) < 1e-9 + assert abs(completion_cost - 0.000375) < 1e-9 + + model_info = litellm.get_model_info(model) + assert model_info["max_input_tokens"] == 260000 + assert model_info["max_output_tokens"] == 65536 + assert model_info["litellm_provider"] == "inception" + assert model_info["mode"] == "chat" + assert model_info["supports_function_calling"] is True + assert model_info["supports_response_schema"] is True diff --git a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py index 86c534c73c2..4c9bd29b337 100644 --- a/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py +++ b/tests/test_litellm/llms/oci/chat/test_oci_chat_transformation.py @@ -1900,3 +1900,198 @@ class TestOCIImageUrlTransformation: adapt_messages_to_generic_oci_standard(messages) assert "image_url" in str(exc_info.value) + + +import itertools +from unittest.mock import patch + +from litellm.llms.oci.chat.transformation import OCIStreamWrapper, _iter_sse_events + +_STREAM_GENERIC_MODEL = "xai.grok-4" +_STREAM_COHERE_MODEL = "cohere.command-latest" + +_GENERIC_TEXT_EVENT = ( + 'data: {{"index":0,"message":{{"role":"ASSISTANT","content":[{{"type":"TEXT","text":"{text}"}}]}},"pad":"aaa"}}' +) +_GENERIC_TERMINAL_EVENT = ( + 'data: {"message":{"role":"ASSISTANT","content":[{"type":"TEXT","text":""}]},"finishReason":"stop","pad":"a"}' +) +_COHERE_TEXT_EVENT = 'data: {{"apiFormat":"COHERE","text":"{text}","pad":"aaaaaa"}}' +_COHERE_TERMINAL_EVENT = ( + 'data: {"apiFormat":"COHERE","text":"123","finishReason":"COMPLETE",' + '"chatHistory":[{"role":"USER","message":"count"},{"role":"CHATBOT","message":"123"}]}' +) + + +def _make_stream_wrapper(model: str) -> OCIStreamWrapper: + logging_obj = MagicMock() + logging_obj.model_call_details = {"custom_llm_provider": "oci", "litellm_params": {}} + return OCIStreamWrapper( + completion_stream=iter([]), + model=model, + custom_llm_provider="oci", + logging_obj=logging_obj, + ) + + +def _ticking_clock(): + """A ``time.time`` stand-in that advances a full second on every call. + + Without it the whole test runs inside one wall-clock second, so a per-chunk + ``created`` would coincidentally match and the drift would go unnoticed. + """ + return itertools.count(1_700_000_000.0) + + +class TestOCIStreamWrapperIdentityPinning: + """One OCI streaming completion must present one id, one created and the + wrapper's model on every chunk, the way every other provider does.""" + + def test_generic_stream_shares_one_id_created_and_model(self): + wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL) + events = [ + _GENERIC_TEXT_EVENT.format(text="1"), + _GENERIC_TEXT_EVENT.format(text="2"), + _GENERIC_TEXT_EVENT.format(text="3"), + _GENERIC_TERMINAL_EVENT, + ] + + with patch("time.time", side_effect=_ticking_clock()): + chunks = [wrapper.chunk_creator(event) for event in events] + + assert len(chunks) == 4 + assert len({chunk.id for chunk in chunks}) == 1 + assert chunks[0].id.startswith("chatcmpl-") + assert len({chunk.created for chunk in chunks}) == 1 + assert {chunk.model for chunk in chunks} == {_STREAM_GENERIC_MODEL} + assert [chunk.choices[0].delta.content for chunk in chunks[:3]] == ["1", "2", "3"] + assert chunks[-1].choices[0].finish_reason == "stop" + assert all(chunk._hidden_params["custom_llm_provider"] == "oci" for chunk in chunks) + + def test_cohere_stream_shares_one_id_created_and_model(self): + """Rebuilding each chunk through the shared creator must not disturb the + Cohere bookkeeping that suppresses the terminal event's repeated text.""" + wrapper = _make_stream_wrapper(_STREAM_COHERE_MODEL) + events = [ + _COHERE_TEXT_EVENT.format(text="1"), + _COHERE_TEXT_EVENT.format(text="2"), + _COHERE_TEXT_EVENT.format(text="3"), + _COHERE_TERMINAL_EVENT, + ] + + with patch("time.time", side_effect=_ticking_clock()): + chunks = [wrapper.chunk_creator(event) for event in events] + + assert len(chunks) == 4 + assert len({chunk.id for chunk in chunks}) == 1 + assert len({chunk.created for chunk in chunks}) == 1 + assert {chunk.model for chunk in chunks} == {_STREAM_COHERE_MODEL} + assert [chunk.choices[0].delta.content for chunk in chunks[:3]] == ["1", "2", "3"] + assert chunks[-1].choices[0].finish_reason == "stop" + assert chunks[-1].choices[0].delta.content is None + assert wrapper._cohere_text_emitted is True + + def test_id_is_pinned_to_the_wrapper_response_id(self): + wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL) + + first = wrapper.chunk_creator(_GENERIC_TEXT_EVENT.format(text="1")) + + assert wrapper.response_id == first.id + assert wrapper.created == first.created + + +class TestOCIStreamWrapperDoneSentinel: + """OCI's GENERIC apiFormat closes the stream with a literal `[DONE]` line; + parsing it as JSON turned every streaming completion into a 500.""" + + @pytest.mark.parametrize("done_event", ["data: [DONE]", "data:[DONE]", "data: [DONE] "]) + def test_done_sentinel_returns_none(self, done_event): + wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL) + assert wrapper.chunk_creator(done_event) is None + + def test_done_sentinel_off_the_sse_splitter_is_skipped(self): + wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL) + wire = ( + f"{_GENERIC_TEXT_EVENT.format(text='1')}\n\n" + f"{_GENERIC_TEXT_EVENT.format(text='2')}\n\n" + f"{_GENERIC_TERMINAL_EVENT}\n\n" + "data: [DONE]\n\n" + ) + + events = list(_iter_sse_events(iter([wire]))) + assert events[-1] == "data: [DONE]" + + chunks = [wrapper.chunk_creator(event) for event in events] + assert chunks[-1] is None + + emitted = [chunk for chunk in chunks if chunk is not None] + assert len(emitted) == 3 + assert len({chunk.id for chunk in emitted}) == 1 + + def test_unparseable_payload_still_raises_oci_error(self): + from litellm.llms.oci.common_utils import OCIError + + wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL) + with pytest.raises(OCIError, match="Chunk cannot be parsed as JSON"): + wrapper.chunk_creator("data: not-json-at-all") + + def test_done_lookalike_payload_still_raises_oci_error(self): + from litellm.llms.oci.common_utils import OCIError + + wrapper = _make_stream_wrapper(_STREAM_GENERIC_MODEL) + with pytest.raises(OCIError, match="Chunk cannot be parsed as JSON"): + wrapper.chunk_creator("data: [DONE] trailing garbage") + + +_GENERIC_TOOL_CALL_EVENT = ( + 'data: {"index":0,"message":{"role":"ASSISTANT","content":[],' + '"toolCalls":[{"type":"FUNCTION","id":"call_1","name":"get_weather","arguments":"{}"}]}}' +) +_GENERIC_TOOL_TERMINAL_EVENT = ( + 'data: {"index":0,"message":{"role":"ASSISTANT","content":[]},"finishReason":"TOOL_CALLS"}' +) + + +def _drain_stream(model: str, events: list[str]) -> list: + logging_obj = MagicMock() + logging_obj.model_call_details = {"custom_llm_provider": "oci", "litellm_params": {}} + wrapper = OCIStreamWrapper( + completion_stream=iter(events), + model=model, + custom_llm_provider="oci", + logging_obj=logging_obj, + ) + return list(wrapper) + + +class TestOCIStreamWrapperTerminalChunk: + """OCI's ``chunk_creator`` override bypasses the shared handler's + finish-reason bookkeeping, so the shared end-of-stream finalizer used to + append a synthetic ``stop`` chunk after OCI's own terminal chunk, silently + downgrading a ``tool_calls`` completion for any client that reads the + finish reason off the last chunk.""" + + def test_generic_tool_call_stream_ends_on_tool_calls(self): + chunks = _drain_stream( + _STREAM_GENERIC_MODEL, + [_GENERIC_TOOL_CALL_EVENT, _GENERIC_TOOL_TERMINAL_EVENT, "data: [DONE]"], + ) + + assert [chunk.choices[0].finish_reason for chunk in chunks] == [None, "tool_calls"] + assert len({chunk.id for chunk in chunks}) == 1 + + def test_generic_text_stream_emits_exactly_one_finish_reason(self): + chunks = _drain_stream( + _STREAM_GENERIC_MODEL, + [_GENERIC_TEXT_EVENT.format(text="1"), _GENERIC_TERMINAL_EVENT, "data: [DONE]"], + ) + + assert [chunk.choices[0].finish_reason for chunk in chunks] == [None, "stop"] + + def test_cohere_stream_emits_exactly_one_finish_reason(self): + chunks = _drain_stream( + _STREAM_COHERE_MODEL, + [_COHERE_TEXT_EVENT.format(text="123"), _COHERE_TERMINAL_EVENT], + ) + + assert [chunk.choices[0].finish_reason for chunk in chunks] == [None, "stop"] diff --git a/tests/test_litellm/llms/reducto/test_parse_v3.py b/tests/test_litellm/llms/reducto/test_parse_v3.py index 140b9737dc0..1d0c826ef8b 100644 --- a/tests/test_litellm/llms/reducto/test_parse_v3.py +++ b/tests/test_litellm/llms/reducto/test_parse_v3.py @@ -1,8 +1,9 @@ import json -import litellm import pytest +import litellm + def _reducto_parse_response() -> dict: return { @@ -68,15 +69,11 @@ def disable_aiohttp_transport(): @pytest.mark.asyncio -async def test_parse_v3_file_upload_and_response_mapping( - disable_aiohttp_transport, respx_mock -): +async def test_parse_v3_file_upload_and_response_mapping(disable_aiohttp_transport, respx_mock): upload_route = respx_mock.post("https://platform.reducto.ai/upload").respond( json={"file_id": "reducto://uploaded.pdf"} ) - parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond( - json=_reducto_parse_response() - ) + parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(json=_reducto_parse_response()) response = await litellm.aocr( model="reducto/parse-v3", @@ -123,15 +120,11 @@ async def test_parse_v3_file_upload_and_response_mapping( @pytest.mark.asyncio -async def test_parse_v3_reducto_id_passthrough_skips_upload( - disable_aiohttp_transport, respx_mock -): +async def test_parse_v3_reducto_id_passthrough_skips_upload(disable_aiohttp_transport, respx_mock): upload_route = respx_mock.post("https://platform.reducto.ai/upload").respond( json={"file_id": "reducto://should-not-upload.pdf"} ) - parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond( - json=_reducto_parse_response() - ) + parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond(json=_reducto_parse_response()) response = await litellm.aocr( model="reducto/parse-v3", @@ -150,3 +143,28 @@ async def test_parse_v3_reducto_id_passthrough_skips_upload( assert parse_request_body["input"] == "reducto://already-uploaded.pdf" assert parse_request_body["retrieval"]["chunk_mode"] == "section" assert response.pages[0].markdown.startswith("Page 1 block A") + + +@pytest.mark.asyncio +async def test_unknown_model_uses_current_protocol_without_local_rejection( + disable_aiohttp_transport, respx_mock +): + parse_route = respx_mock.post("https://platform.reducto.ai/parse").respond( + json=_reducto_parse_response() + ) + + response = await litellm.aocr( + model="reducto/future-parse-model", + document={ + "type": "document_url", + "document_url": "reducto://already-uploaded.pdf", + }, + api_key="test-key", + api_base="https://platform.reducto.ai", + ) + + assert parse_route.called + assert json.loads(parse_route.calls[0].request.read()) == { + "input": "reducto://already-uploaded.pdf" + } + assert response.model == "future-parse-model" diff --git a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py index d2788408e09..101f6e6fa5d 100644 --- a/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py +++ b/tests/test_litellm/llms/vertex_ai/gemini/test_vertex_and_google_ai_studio_gemini.py @@ -5769,3 +5769,70 @@ def test_calculate_web_search_requests_counts_unique_queries(): assert VertexGeminiConfig._calculate_web_search_requests([]) is None assert VertexGeminiConfig._calculate_web_search_requests([{"webSearchQueries": ["", ""]}]) is None + + +@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai"]) +@pytest.mark.parametrize( + "model", + ["gemini-2.5-flash", "gemini-3-pro-preview"], + ids=["thinking_budget_mapper", "thinking_level_mapper"], +) +@pytest.mark.parametrize("reasoning_effort", ["banana", "xhigh"]) +def test_invalid_reasoning_effort_is_a_400_not_a_500(custom_llm_provider, model, reasoning_effort): + """Regression for #40474. + + Both reasoning_effort mappers used to end their if/elif chain in a bare `ValueError`, which + `exception_type()` has no branch for, so it fell through to `APIConnectionError` and the proxy + answered a malformed client request with a retryable HTTP 500. `xhigh` is covered alongside the + nonsense value because it is a member of litellm's own `REASONING_EFFORT` literal, so callers + bridging from OpenAI-shaped code reach it without typing anything wrong. + """ + from litellm.utils import get_optional_params + + with pytest.raises(litellm.BadRequestError) as exc_info: + get_optional_params( + model=model, + custom_llm_provider=custom_llm_provider, + reasoning_effort=reasoning_effort, + drop_params=True, + ) + + assert exc_info.value.status_code == 400 + message: Final = str(exc_info.value) + assert reasoning_effort in message + for supported in ("minimal", "low", "medium", "high", "none", "disable"): + assert supported in message + + +@pytest.mark.parametrize("custom_llm_provider", ["gemini", "vertex_ai"]) +def test_invalid_reasoning_effort_surfaces_as_400_through_completion(custom_llm_provider): + """The same request through `completion()` must not come back as a retryable 500. + + Needs no provider credentials: param mapping runs before any network call. + """ + with pytest.raises(litellm.BadRequestError) as exc_info: + completion( + model=f"{custom_llm_provider}/gemini-3-pro-preview", + messages=[{"role": "user", "content": "hi"}], + reasoning_effort="banana", + ) + + assert exc_info.value.status_code == 400 + assert not isinstance(exc_info.value, litellm.APIConnectionError) + + +@pytest.mark.parametrize("model", ["gemini-2.5-flash", "gemini-3-pro-preview"]) +def test_supported_reasoning_efforts_still_map(model): + """Guards the fix against over-rejecting: every advertised value must still produce a config.""" + from litellm.llms.vertex_ai.gemini.vertex_and_google_ai_studio_gemini import ( + SUPPORTED_REASONING_EFFORTS, + ) + + for effort in SUPPORTED_REASONING_EFFORTS: + result: Final = VertexGeminiConfig().map_openai_params( + non_default_params={"reasoning_effort": effort}, + optional_params={}, + model=model, + drop_params=False, + ) + assert "thinkingConfig" in result diff --git a/tests/test_litellm/models/test_models.py b/tests/test_litellm/models/test_models.py index 9ae9b732066..aa6449c98dd 100644 --- a/tests/test_litellm/models/test_models.py +++ b/tests/test_litellm/models/test_models.py @@ -8,6 +8,7 @@ import pytest from pydantic import BaseModel, TypeAdapter from litellm.models.access_group import LiteLLM_AccessGroupTable +from litellm.models.autorouter_session import LiteLLM_AutoRouterSession from litellm.models.budget import ( LiteLLM_BudgetTable, LiteLLM_BudgetTableFull, @@ -588,3 +589,35 @@ class TestManagedTables: ) assert table.vector_store_id == "vs1" assert table.custom_llm_provider == "openai" + + +class TestAutoRouterSession: + @staticmethod + def _row(baseline_models: dict) -> LiteLLM_AutoRouterSession: + return LiteLLM_AutoRouterSession( + api_key="k", + session_id="s", + router_name="auto", + router_type="complexity", + first_turn_at=datetime(2026, 9, 1, 12, 0, 0), + last_turn_at=datetime(2026, 9, 1, 12, 5, 0), + last_model="anthropic/claude-sonnet-5", + turns=3, + spend=0.14, + saved_spend=0.24, + classifier_cost=0.0, + tier_turns={}, + baseline_models=baseline_models, + ) + + def test_the_baseline_label_is_the_one_most_turns_were_priced_against(self): + assert self._row({"anthropic/claude-opus-5": 2, "anthropic/claude-sonnet-5": 1}).baseline_model == ( + "anthropic/claude-opus-5" + ) + + def test_a_tie_between_baselines_is_broken_deterministically(self): + assert self._row({"b-model": 1, "a-model": 1}).baseline_model == "b-model" + assert self._row({"a-model": 1, "b-model": 1}).baseline_model == "b-model" + + def test_a_row_whose_turns_recorded_no_baseline_has_no_label(self): + assert self._row({}).baseline_model is None diff --git a/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py b/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py index 0c8b1cc2836..460aff3e8d1 100644 --- a/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py +++ b/tests/test_litellm/ocr/test_ocr_azure_document_intelligence_api_base.py @@ -1,20 +1,18 @@ """ -Regression tests for Azure Document Intelligence api_base resolution in OCR. +Regression tests for Azure Document Intelligence api_base ownership in OCR. `azure_ai` exposes two OCR services on one provider; the `doc-intelligence` -sub-route must resolve to `AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT`, not to the -generic `AZURE_AI_API_BASE` fallback that `get_llm_provider` injects. These tests -pin that routing and guard the backwards-compatibility contract that an explicitly -supplied api_base is always honoured. +sub-route must defer environment resolution to Rust, not accept the generic +`AZURE_AI_API_BASE` fallback that `get_llm_provider` injects. An explicitly +supplied api_base is still always honoured. """ from litellm.llms.azure_ai.ocr.common_utils import ( is_azure_document_intelligence_model, ) -from litellm.ocr.main import _prepare_ocr_request, _rust_bridge_api_base +from litellm.ocr.main import _prepare_ocr_request _DOC = {"type": "document_url", "document_url": "https://example.com/doc.pdf"} -_DOC_INTELLIGENCE_ENDPOINT = "https://di.cognitiveservices.azure.com" _AZURE_AI_API_BASE = "https://generic-azure-ai.example.com" @@ -23,13 +21,6 @@ class _FakeLogging: return None -def _resolve_secret(name: str) -> str | None: - return { - "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT": _DOC_INTELLIGENCE_ENDPOINT, - "AZURE_AI_API_BASE": _AZURE_AI_API_BASE, - }.get(name) - - def _prepare(model: str, api_base: str | None): return _prepare_ocr_request( model=model, @@ -56,15 +47,13 @@ class TestIsAzureDocumentIntelligenceModel: class TestDocIntelligenceApiBaseResolution: def test_generic_azure_ai_base_does_not_hijack_doc_intelligence(self, monkeypatch): - """Without an explicit api_base, the AZURE_AI_API_BASE fallback must not - overwrite the endpoint, so it resolves to the Document Intelligence one.""" + """The generic Azure base must not overwrite Rust-owned DI resolution.""" monkeypatch.setenv("AZURE_AI_API_BASE", _AZURE_AI_API_BASE) monkeypatch.delenv("AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT", raising=False) prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", None) assert prepared.api_base is None - assert _rust_bridge_api_base(prepared, _resolve_secret) == _DOC_INTELLIGENCE_ENDPOINT def test_explicit_api_base_is_honoured_for_doc_intelligence(self, monkeypatch): """A caller-supplied api_base must always win, even for doc-intelligence.""" @@ -74,7 +63,6 @@ class TestDocIntelligenceApiBaseResolution: prepared = _prepare("azure_ai/doc-intelligence/prebuilt-layout", custom) assert prepared.api_base == custom - assert _rust_bridge_api_base(prepared, _resolve_secret) == custom def test_generic_azure_ai_base_still_applies_to_mistral_ocr(self, monkeypatch): """Non doc-intelligence azure_ai models keep using AZURE_AI_API_BASE.""" diff --git a/tests/test_litellm/ocr/test_ocr_native_format.py b/tests/test_litellm/ocr/test_ocr_native_format.py index 249fbda713e..46e9a4d3729 100644 --- a/tests/test_litellm/ocr/test_ocr_native_format.py +++ b/tests/test_litellm/ocr/test_ocr_native_format.py @@ -1,52 +1,60 @@ """ -Tests for the OCR `req_format` option in the SDK request path: -providers that don't support a native response must reject it, and the Rust -bridge (which only returns the normalized shape) must not serve native requests. +Tests for the OCR `req_format` option in the SDK request path. """ -import dataclasses -from unittest.mock import MagicMock - import pytest import litellm -from litellm.llms.azure_ai.ocr.cohere_parse_transformation import AzureAICohereParseConfig -from litellm.llms.cohere.ocr.transformation import CohereParseConfig -from litellm.ocr.main import _PreparedOCRRequest, _rust_ocr_supported +from litellm.rust_bridge import ocr as rust_ocr_bridge +from litellm.rust_bridge.ocr import LiteLLMOcrRequest DOCUMENT = {"type": "document_url", "document_url": "https://example.com/doc.pdf"} -def _prepared(optional_params: dict[str, object]) -> _PreparedOCRRequest: - return _PreparedOCRRequest( - model="doc-intelligence/prebuilt-layout", - document=dict(DOCUMENT), +def _request( + optional_params: dict[str, object], model: str = "azure_ai/doc-intelligence/prebuilt-layout" +) -> LiteLLMOcrRequest: + return LiteLLMOcrRequest( + model=model, + document=DOCUMENT, api_key="fake-key", - api_base="https://example.cognitiveservices.azure.com", - custom_llm_provider="azure_ai", + api_base=None, + custom_llm_provider=None, extra_headers=None, - provider_config=MagicMock(), - optional_params=optional_params, - litellm_params={}, - effective_timeout=60.0, - litellm_logging_obj=MagicMock(), + timeout=60.0, + kwargs=optional_params, ) @pytest.mark.parametrize("optional_params", [{}, {"req_format": "litellm"}]) def test_rust_ocr_serves_default_format(optional_params): - assert _rust_ocr_supported(_prepared(optional_params)) is True + assert rust_ocr_bridge.supported(_request(optional_params)) is True -def test_rust_ocr_skipped_for_native_format(): - assert _rust_ocr_supported(_prepared({"req_format": "native"})) is False +def test_rust_ocr_serves_native_format_for_document_intelligence(): + assert rust_ocr_bridge.supported(_request({"req_format": "native"})) is True -@pytest.mark.parametrize("provider_config", [CohereParseConfig(), AzureAICohereParseConfig()]) -def test_rust_ocr_skipped_for_configs_without_bridge_support(provider_config): - prepared = dataclasses.replace(_prepared({}), provider_config=provider_config) +def test_rust_ocr_response_retains_provider_native_response(): + provider_response = {"status": "succeeded", "analyzeResult": {"content": "native"}} + response = rust_ocr_bridge._response( + { + "pages": [], + "model": "prebuilt-layout", + "document_annotation": None, + "usage_info": {"pages_processed": 0}, + "object": "ocr", + "provider_native_response": provider_response, + } + ) - assert _rust_ocr_supported(prepared) is False + assert response.get_provider_native_response() == provider_response + assert response.model_dump().get("provider_native_response") is None + + +@pytest.mark.parametrize("model", ["cohere/cohere-parse", "azure_ai/cohere-parse"]) +def test_rust_ocr_skipped_for_unsupported_models(model): + assert rust_ocr_bridge.supported(_request({}, model)) is False @pytest.mark.asyncio diff --git a/tests/test_litellm/ocr/test_rust_bridge.py b/tests/test_litellm/ocr/test_rust_bridge.py index c34833221cc..dbb4f822d0b 100644 --- a/tests/test_litellm/ocr/test_rust_bridge.py +++ b/tests/test_litellm/ocr/test_rust_bridge.py @@ -3,7 +3,6 @@ import builtins import importlib import types -from typing import Any import httpx import pytest @@ -59,6 +58,7 @@ class RecordingBridge: custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], + input_sources: dict[str, str], timeout_seconds: float | None, ) -> dict[str, object]: self.calls.append( @@ -70,6 +70,7 @@ class RecordingBridge: "custom_llm_provider": custom_llm_provider, "extra_headers": extra_headers, "optional_params": optional_params, + "input_sources": input_sources, "timeout_seconds": timeout_seconds, } ) @@ -91,6 +92,7 @@ class RecordingAsyncBridge: custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], + input_sources: dict[str, str], timeout_seconds: float | None, ) -> dict[str, object]: self.calls.append( @@ -102,6 +104,7 @@ class RecordingAsyncBridge: "custom_llm_provider": custom_llm_provider, "extra_headers": extra_headers, "optional_params": optional_params, + "input_sources": input_sources, "timeout_seconds": timeout_seconds, } ) @@ -118,6 +121,7 @@ class RaisingBridge: custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], + input_sources: dict[str, str], timeout_seconds: float | None, ) -> dict[str, object]: raise RuntimeError("bridge failed") @@ -133,6 +137,7 @@ class RaisingAsyncBridge: custom_llm_provider: str | None, extra_headers: dict[str, object] | None, optional_params: dict[str, object], + input_sources: dict[str, str], timeout_seconds: float | None, ) -> dict[str, object]: raise RuntimeError("bridge failed") @@ -144,6 +149,9 @@ class RecordingLogging: def __init__(self) -> None: self.pre_call_kwargs: dict[str, object] | None = None + def update_from_kwargs(self, **kwargs: object) -> None: + self.update_kwargs = kwargs + def pre_call( self, *, @@ -158,66 +166,32 @@ class RecordingLogging: } -class FakeOCRConfig: - """A stand-in ``BaseOCRConfig`` that echoes the request it would build.""" - - def __init__(self, api_key_env_var: str = "MISTRAL_API_KEY") -> None: - self.api_key_env_var = api_key_env_var - - def get_api_key_env_var(self) -> str: - return self.api_key_env_var - - def validate_environment( - self, - *, - headers: dict[str, object], - model: str, - api_key: str | None, - api_base: str | None, - litellm_params: dict[str, object], - ) -> dict[str, object]: - return {"Authorization": f"Bearer {api_key}", **headers} - - def get_complete_url( - self, - *, - api_base: str | None, - model: str, - optional_params: dict[str, object], - litellm_params: dict[str, object], - ) -> str: - return f"{api_base or 'https://api.mistral.ai/v1'}/ocr" - - def get_error_class(self, error_message: str, status_code: int, headers: dict[str, str]) -> BaseLLMException: - return BaseLLMException(status_code=status_code, message=error_message, headers=headers) - - -def build_prepared_request( +def build_request( *, logging_obj: RecordingLogging | None = None, - provider_config: FakeOCRConfig | None = None, model: str = "mistral-ocr-latest", document: dict[str, object] = DOCUMENT, api_key: str | None = "sk-test", api_base: str | None = None, - custom_llm_provider: str = "mistral", + custom_llm_provider: str | None = "mistral", extra_headers: dict[str, object] | None = None, optional_params: dict[str, object] | None = None, litellm_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = 12.5, -) -> Any: - return ocr_main._PreparedOCRRequest( +) -> rust_bridge.LiteLLMOcrRequest: + return rust_bridge.LiteLLMOcrRequest( model=model, document=document, api_key=api_key, api_base=api_base, custom_llm_provider=custom_llm_provider, extra_headers=extra_headers, - provider_config=provider_config or FakeOCRConfig(), - optional_params=optional_params or {}, - litellm_params=litellm_params or {}, - effective_timeout=timeout, - litellm_logging_obj=logging_obj or RecordingLogging(), + timeout=timeout, + kwargs={ + **(optional_params or {}), + **(litellm_params or {}), + "litellm_logging_obj": logging_obj or RecordingLogging(), + }, ) @@ -425,6 +399,7 @@ def test_bridge_wrapper_forwards_prepared_args_and_wraps_response(): "x-trace-id": "trace-1", }, "optional_params": {"include_image_base64": True, "pages": [0]}, + "input_sources": {}, "timeout_seconds": 12.5, } @@ -456,6 +431,7 @@ async def test_bridge_wrapper_forwards_prepared_async_args_and_wraps_response(): "custom_llm_provider": "vertex_ai", "extra_headers": None, "optional_params": {"vertex_project": "project-1"}, + "input_sources": {}, "timeout_seconds": 42.0, } @@ -467,7 +443,7 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): rust_bridge._OCR.override(bridge) response = ocr_main._run_rust_ocr( - prepared_request=build_prepared_request( + request=build_request( logging_obj=logging_obj, api_base="https://proxy.internal", extra_headers={"x-trace-id": "trace-1"}, @@ -486,10 +462,10 @@ def test_run_rust_ocr_prepares_request_and_wraps_response(): "api_base": "https://proxy.internal", "custom_llm_provider": "mistral", "extra_headers": { - "Authorization": "Bearer sk-test", "x-trace-id": "trace-1", }, "optional_params": {"include_image_base64": True}, + "input_sources": {}, "timeout_seconds": 12.5, } @@ -499,7 +475,7 @@ def test_rust_upstream_error_uses_ocr_provider_error_mapping(): mapped = ocr_main._map_rust_ocr_error( error, - build_prepared_request(), + build_request(), (RuntimeError, RustUpstreamError), ) @@ -514,7 +490,7 @@ def test_run_rust_ocr_resolves_key_via_secret_manager_when_missing(): rust_bridge._OCR.override(bridge) ocr_main._run_rust_ocr( - prepared_request=build_prepared_request(api_key=None, timeout=None), + request=build_request(api_key=None, timeout=None), resolve_api_key=lambda name: "sk-from-vault" if name == "MISTRAL_API_KEY" else None, ) @@ -530,7 +506,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver(): raise AssertionError(f"resolver should not be called for {name}") ocr_main._run_rust_ocr( - prepared_request=build_prepared_request( + request=build_request( api_key="sk-explicit", timeout=None, ), @@ -540,7 +516,7 @@ def test_run_rust_ocr_prefers_explicit_key_over_resolver(): assert bridge.calls[0]["api_key"] == "sk-explicit" -def test_run_rust_ocr_uses_provider_api_key_env_var(): +def test_run_rust_ocr_uses_mistral_secret_manager_without_provider_config(): bridge = RecordingBridge() resolver_calls = [] litellm.rust(True) @@ -551,16 +527,15 @@ def test_run_rust_ocr_uses_provider_api_key_env_var(): return "sk-provider-env" ocr_main._run_rust_ocr( - prepared_request=build_prepared_request( - provider_config=FakeOCRConfig(api_key_env_var="PROVIDER_OCR_API_KEY"), - model="provider-ocr-model", + request=build_request( + model="mistral-ocr-latest", api_key=None, timeout=None, ), resolve_api_key=_resolver, ) - assert resolver_calls == ["PROVIDER_OCR_API_KEY"] + assert resolver_calls == ["MISTRAL_API_KEY"] assert bridge.calls[0]["api_key"] == "sk-provider-env" @@ -570,7 +545,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): rust_bridge._OCR.override(bridge) ocr_main._run_rust_ocr( - prepared_request=build_prepared_request( + request=build_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", litellm_params={ @@ -588,6 +563,7 @@ def test_prepare_rust_ocr_call_forwards_vertex_routing_metadata(): "include_image_base64": True, "vertex_project": "project-1", "vertex_location": "us-central1", + "vertex_credentials": "redacted", } @@ -600,10 +576,11 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana return { "VERTEXAI_PROJECT": "project-from-secret", "VERTEXAI_LOCATION": "us-east5", + "VERTEXAI_CREDENTIALS": "credentials-from-secret", }.get(name) ocr_main._run_rust_ocr( - prepared_request=build_prepared_request( + request=build_request( custom_llm_provider="vertex_ai", model="mistral-ocr-maas", timeout=None, @@ -613,44 +590,210 @@ def test_prepare_rust_ocr_call_resolves_vertex_routing_metadata_from_secret_mana assert bridge.calls[0]["optional_params"]["vertex_project"] == "project-from-secret" assert bridge.calls[0]["optional_params"]["vertex_location"] == "us-east5" + assert bridge.calls[0]["optional_params"]["vertex_credentials"] == "credentials-from-secret" -def test_prepare_rust_ocr_call_resolves_azure_ai_api_base_from_secret_manager(): +def test_prepare_rust_ocr_call_defers_azure_environment_resolution_to_rust(): bridge = RecordingBridge() litellm.rust(True) rust_bridge._OCR.override(bridge) ocr_main._run_rust_ocr( - prepared_request=build_prepared_request( + request=build_request( custom_llm_provider="azure_ai", model="pixtral-12b-2409", + api_key=None, api_base=None, timeout=None, ), - resolve_api_key=lambda name: "https://azure.example.com" if name == "AZURE_AI_API_BASE" else None, + resolve_api_key=lambda name: pytest.fail(f"Python resolved Azure secret {name}"), ) - assert bridge.calls[0]["api_base"] == "https://azure.example.com" + assert bridge.calls[0]["api_base"] is None + assert bridge.calls[0]["api_key"] is None + assert bridge.calls[0]["extra_headers"] is None -def test_prepare_rust_ocr_call_resolves_document_intelligence_endpoint(): +def test_prepare_rust_ocr_call_defers_document_intelligence_environment_to_rust(): bridge = RecordingBridge() litellm.rust(True) rust_bridge._OCR.override(bridge) ocr_main._run_rust_ocr( - prepared_request=build_prepared_request( + request=build_request( custom_llm_provider="azure_ai", model="doc-intelligence/prebuilt-layout", api_base=None, timeout=None, ), - resolve_api_key=lambda name: ( - "https://document-intelligence.example.com" if name == "AZURE_DOCUMENT_INTELLIGENCE_ENDPOINT" else None - ), + resolve_api_key=lambda name: pytest.fail(f"Python resolved Azure secret {name}"), ) - assert bridge.calls[0]["api_base"] == "https://document-intelligence.example.com" + assert bridge.calls[0]["api_base"] is None + + +def test_prepare_rust_ocr_call_forwards_raw_azure_auth_inputs(): + bridge = RecordingBridge() + litellm.rust(True) + rust_bridge._OCR.override(bridge) + + ocr_main._run_rust_ocr( + request=build_request( + custom_llm_provider="azure_ai", + model="pixtral-12b-2409", + api_key=None, + api_base="https://azure.example.com", + extra_headers={"x-trace-id": "trace-1"}, + litellm_params={ + "azure_ad_token": "entra-token", + "tenant_id": "tenant", + "client_id": "client", + "client_secret": "secret", + "azure_scope": "scope", + "azure_authority_host": "https://login.example.com", + "azure_credential": "ClientSecretCredential", + "azure_federated_token_file": "/token", + }, + timeout=None, + ), + resolve_api_key=lambda name: pytest.fail(f"Python resolved Azure secret {name}"), + ) + + call = bridge.calls[0] + assert call["api_key"] is None + assert call["api_base"] == "https://azure.example.com" + assert call["extra_headers"] == {"x-trace-id": "trace-1"} + assert call["optional_params"] == { + "azure_ad_token": "entra-token", + "tenant_id": "tenant", + "client_id": "client", + "client_secret": "secret", + "azure_scope": "scope", + "azure_authority_host": "https://login.example.com", + "azure_credential": "ClientSecretCredential", + "azure_federated_token_file": "/token", + } + assert call["input_sources"] == {} + + +def test_prepare_rust_ocr_call_preserves_proxy_input_sources(): + bridge = RecordingBridge() + litellm.rust(True) + rust_bridge._OCR.override(bridge) + request_values = { + "tenant_id": "tenant", + "client_id": "client", + "client_secret": "secret", + "azure_authority_host": "https://login.example.com", + "api_base": "https://azure.example.com", + } + + ocr_main._run_rust_ocr( + request=build_request( + custom_llm_provider="azure_ai", + model="pixtral-12b-2409", + api_key="request-key", + api_base="https://azure.example.com", + litellm_params={ + "tenant_id": "tenant", + "client_id": "client", + "client_secret": "secret", + "azure_authority_host": "https://login.example.com", + "proxy_server_request": {"body": request_values, "credential_fields": ("api_key",)}, + }, + ), + resolve_api_key=lambda _name: None, + ) + + assert bridge.calls[0]["input_sources"] == { + **{name: "request" for name in request_values}, + "api_key": "request", + } + + marshaled = rust_bridge._marshal( + build_request( + custom_llm_provider="azure_ai", + model="pixtral-12b-2409", + api_key="request-key", + api_base="https://azure.example.com", + litellm_params={ + "proxy_server_request": { + "body": {"api_base": "https://azure.example.com"}, + "credential_fields": ("api_key",), + } + }, + ), + lambda _name: None, + lambda document: document, + ) + assert marshaled.input_sources == {"api_base": "request", "api_key": "request"} + + +def test_rust_ocr_logging_redacts_azure_credentials(): + bridge = RecordingBridge() + logging_obj = RecordingLogging() + litellm.rust(True) + rust_bridge._OCR.override(bridge) + + ocr_main._run_rust_ocr( + request=build_request( + logging_obj=logging_obj, + custom_llm_provider="azure_ai", + model="pixtral-12b-2409", + api_key=None, + litellm_params={"azure_ad_token": "token", "client_secret": "secret"}, + ), + resolve_api_key=lambda _name: None, + ) + + assert logging_obj.update_kwargs["optional_params"] == { + "azure_ad_token": "****", + "client_secret": "****", + } + assert logging_obj.pre_call_kwargs is not None + additional_args = logging_obj.pre_call_kwargs["additional_args"] + assert isinstance(additional_args, dict) + complete_input = additional_args["complete_input_dict"] + assert isinstance(complete_input, dict) + assert complete_input["azure_ad_token"] == "****" + assert complete_input["client_secret"] == "****" + + +def test_rust_eligibility_rejects_python_only_azure_auth_modes(): + for params in ( + {"azure_ad_token_provider": lambda: "token"}, + {"azure_username": "user"}, + {"azure_password": "password"}, + ): + assert not ocr_main._rust_ocr_supported( + build_request( + custom_llm_provider="azure_ai", + model="pixtral-12b-2409", + litellm_params=params, + ) + ) + + +def test_prepare_rust_ocr_call_forwards_global_azure_refresh(monkeypatch: pytest.MonkeyPatch): + bridge = RecordingBridge() + litellm.rust(True) + rust_bridge._OCR.override(bridge) + monkeypatch.setattr(litellm, "enable_azure_ad_token_refresh", True) + + ocr_main._run_rust_ocr( + request=build_request( + custom_llm_provider="azure_ai", + model="pixtral-12b-2409", + api_key=None, + api_base="https://azure.example.com", + litellm_params={"proxy_server_request": {"body": {"enable_azure_ad_token_refresh": True}}}, + timeout=None, + ), + resolve_api_key=lambda _name: None, + ) + + assert bridge.calls[0]["optional_params"] == {"enable_azure_ad_token_refresh": True} + assert bridge.calls[0]["input_sources"] == {"enable_azure_ad_token_refresh": "deployment"} def test_run_rust_ocr_runs_pre_call_logging(): @@ -660,7 +803,7 @@ def test_run_rust_ocr_runs_pre_call_logging(): rust_bridge._OCR.override(bridge) ocr_main._run_rust_ocr( - prepared_request=build_prepared_request( + request=build_request( logging_obj=logging_obj, api_base="https://api.mistral.ai/v1", extra_headers={"x-trace-id": "trace-1"}, @@ -676,9 +819,8 @@ def test_run_rust_ocr_runs_pre_call_logging(): complete_input = additional_args["complete_input_dict"] assert complete_input["document"] == DOCUMENT assert complete_input["include_image_base64"] is True - assert additional_args["api_base"] == "https://api.mistral.ai/v1/ocr" + assert additional_args["api_base"] == "https://api.mistral.ai/v1" assert additional_args["headers"] == { - "Authorization": "Bearer sk-test", "x-trace-id": "trace-1", } @@ -696,12 +838,11 @@ def test_ocr_routes_to_rust_when_enabled(fake_bridge): assert response.pages[0].markdown == "hello world" assert len(fake_bridge.calls) == 1 call = fake_bridge.calls[0] - assert call["model"] == "mistral-ocr-latest" + assert call["model"] == MODEL assert call["document"] == DOCUMENT assert call["api_key"] == "sk-test" - assert call["custom_llm_provider"] == "mistral" + assert call["custom_llm_provider"] is None assert call["extra_headers"] == { - "Authorization": "Bearer sk-test", "x-trace-id": "trace-1", } assert call["optional_params"].get("include_image_base64") is True @@ -717,8 +858,29 @@ def test_ocr_routes_azure_ai_to_rust_when_enabled(fake_bridge): assert isinstance(response, OCRResponse) assert len(fake_bridge.calls) == 1 - assert fake_bridge.calls[0]["model"] == "pixtral-12b-2409" - assert fake_bridge.calls[0]["custom_llm_provider"] == "azure_ai" + assert fake_bridge.calls[0]["model"] == "azure_ai/pixtral-12b-2409" + assert fake_bridge.calls[0]["custom_llm_provider"] is None + assert fake_bridge.calls[0]["extra_headers"] is None + + +def test_ocr_routes_azure_entra_inputs_to_rust_without_python_auth(fake_bridge): + response = litellm.ocr( + model="azure_ai/pixtral-12b-2409", + document=DOCUMENT, + api_base="https://example.services.ai.azure.com", + azure_ad_token="entra-token", + tenant_id="tenant", + client_id="client", + ) + + assert isinstance(response, OCRResponse) + assert fake_bridge.calls[0]["api_key"] is None + assert fake_bridge.calls[0]["extra_headers"] is None + assert fake_bridge.calls[0]["optional_params"] == { + "azure_ad_token": "entra-token", + "tenant_id": "tenant", + "client_id": "client", + } def test_ocr_rust_path_converts_file_document_before_bridge(fake_bridge): @@ -768,12 +930,11 @@ async def test_aocr_routes_to_async_rust_when_enabled(fake_async_bridge): assert response.pages[0].markdown == "hello world" assert len(fake_async_bridge.calls) == 1 call = fake_async_bridge.calls[0] - assert call["model"] == "mistral-ocr-latest" + assert call["model"] == MODEL assert call["document"] == DOCUMENT assert call["api_key"] == "sk-test" - assert call["custom_llm_provider"] == "mistral" + assert call["custom_llm_provider"] is None assert call["extra_headers"] == { - "Authorization": "Bearer sk-test", "x-trace-id": "trace-1", } assert call["optional_params"].get("include_image_base64") is True @@ -864,3 +1025,137 @@ def test_ocr_provider_configs_expose_api_key_env_vars(): assert AzureDocumentIntelligenceOCRConfig().get_api_key_env_var() == "AZURE_DOCUMENT_INTELLIGENCE_API_KEY" assert VertexAIOCRConfig().get_api_key_env_var() == "VERTEX_AI_API_KEY" assert VertexAIDeepSeekOCRConfig().get_api_key_env_var() == "VERTEX_AI_API_KEY" + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.asyncio +async def test_rust_receives_unmapped_azure_options(asynchronous, fake_bridge, fake_async_bridge): + from typing import Final + + arguments: Final = { + "model": "azure_ai/doc-intelligence/prebuilt-layout", + "document": DOCUMENT, + "api_key": "test-key", + "pages": [0, 2], + "features": ["languages", "style"], + "provider_extension": {"enabled": True}, + } + if asynchronous: + await litellm.aocr(**arguments) + else: + litellm.ocr(**arguments) + call: Final = (fake_async_bridge if asynchronous else fake_bridge).calls[0] + assert call["model"] == arguments["model"] + assert call["custom_llm_provider"] is None + assert call["extra_headers"] is None + assert call["optional_params"] == { + "pages": [0, 2], + "features": ["languages", "style"], + "provider_extension": {"enabled": True}, + } + + +@pytest.mark.parametrize("enabled", [False, True]) +@pytest.mark.asyncio +async def test_python_fallback_maps_original_options_once(enabled, monkeypatch): + from io import BytesIO + from typing import Final + + class PythonHandler: + def __init__(self): + self.calls = [] + + def ocr(self, **kwargs): + self.calls.append(kwargs) + return OCRResponse(pages=[], model=kwargs["model"]) + + handler: Final = PythonHandler() + monkeypatch.setattr(ocr_main, "base_llm_http_handler", handler) + litellm.rust(enabled) + rust_bridge._OCR.override(None) + rust_bridge._AOCR.override(None) + for asynchronous in (False, True): + file: Final = BytesIO(b"test document") + arguments: Final = { + "model": "azure_ai/doc-intelligence/prebuilt-layout", + "document": {"type": "file", "file": file}, + "api_key": "test-key", + "pages": [0, 2], + } + if asynchronous: + await litellm.aocr(**arguments) + else: + litellm.ocr(**arguments) + assert handler.calls[-1]["optional_params"]["pages"] == "1,3" + assert handler.calls[-1]["document"]["document_url"].endswith("dGVzdCBkb2N1bWVudA==") + assert len(handler.calls) == 2 + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/doc-intelligence/prebuilt-read"]) +@pytest.mark.asyncio +async def test_native_public_ocr_matches_python(model, asynchronous): + import json + from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + from threading import Thread + from typing import Final + from urllib.parse import parse_qsl, urlsplit + + native: Final = rust_bridge_loader.get_native_bridge() + if native is None: + pytest.skip("requires the compiled Rust extension") + calls: Final = [] + + class Handler(BaseHTTPRequestHandler): + def do_POST(self): + body: Final = json.loads(self.rfile.read(int(self.headers["Content-Length"]))) + target: Final = urlsplit(self.path) + calls.append( + ( + target.path, + parse_qsl(target.query), + self.headers.get("Authorization"), + self.headers.get("Ocp-Apim-Subscription-Key"), + body, + ) + ) + payload: Final = ( + {"status": "succeeded", "analyzeResult": {"pages": []}} + if "doc-intelligence" in model + else {"pages": [{"index": 0, "markdown": "hello"}]} + ) + encoded: Final = json.dumps(payload).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(encoded))) + self.end_headers() + self.wfile.write(encoded) + + def log_message(self, *_args): + pass + + server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + thread: Final = Thread(target=server.serve_forever, daemon=True) + thread.start() + responses: Final = [] + try: + for enabled in (False, True): + litellm.rust(enabled) + arguments: Final = { + "model": model, + "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + "api_key": "test-key", + "api_base": f"http://127.0.0.1:{server.server_port}", + "pages": [0, 2], + "timeout": 3.0, + } + response: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments) + responses.append(response.model_dump()) + assert len(calls) == 2 + assert calls[0] == calls[1] + for key in ("model", "pages", "object"): + assert responses[0][key] == responses[1][key] + finally: + server.shutdown() + server.server_close() + thread.join(timeout=3) diff --git a/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py b/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py index a88b0ef0c4b..5e13db9439b 100644 --- a/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py +++ b/tests/test_litellm/passthrough/test_streaming_interrupt_spend_tracking.py @@ -34,6 +34,34 @@ class _ImmediateExecutor: fn(*args, **kwargs) +class _RecordingCollector: + def __init__(self) -> None: + self.chunks: List[bytes] = [] + + def add(self, chunk: bytes) -> None: + self.chunks.append(chunk) + + def build_logged_response(self, litellm_logging_obj: MagicMock) -> bytes: + return b"".join(self.chunks) + + +class _FailingCollector(_RecordingCollector): + def add(self, chunk: bytes) -> None: + raise ValueError("bad frame") + + +def _provider_config(collector: _RecordingCollector) -> MagicMock: + provider_config = MagicMock() + provider_config.create_stream_collector.return_value = collector + return provider_config + + +def _spend_payload(flush_mock: MagicMock) -> bytes: + flush_mock.assert_called_once() + collector = flush_mock.call_args.kwargs["collector"] + return collector.build_logged_response(litellm_logging_obj=MagicMock()) + + @pytest.mark.asyncio async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion(): from litellm.passthrough.main import AsyncPassthroughStreamingResponse @@ -48,13 +76,12 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion(): return mock_response mock_logging_obj = _make_logging_obj() - provider_config = MagicMock() received = [] received_response = AsyncPassthroughStreamingResponse( response=response_coro(), litellm_logging_obj=mock_logging_obj, - provider_config=provider_config, + provider_config=_provider_config(_RecordingCollector()), ) async for chunk in received_response: @@ -67,12 +94,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_normal_completion(): await asyncio.sleep(0) - mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once() - call_kwargs = ( - mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs - ) - assert call_kwargs["raw_bytes"] == chunks - assert call_kwargs["provider_config"] is provider_config + assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == b"".join(chunks) @pytest.mark.asyncio @@ -93,12 +115,11 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_client_disconnect(): return mock_response mock_logging_obj = _make_logging_obj() - provider_config = MagicMock() gen = AsyncPassthroughStreamingResponse( response=response_coro(), litellm_logging_obj=mock_logging_obj, - provider_config=provider_config, + provider_config=_provider_config(_RecordingCollector()), ) received = [await gen.__anext__()] @@ -108,11 +129,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_client_disconnect(): await asyncio.sleep(0) - mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once() - call_kwargs = ( - mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs - ) - assert call_kwargs["raw_bytes"] == [chunks[0]] + assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == chunks[0] @pytest.mark.asyncio @@ -178,14 +195,13 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_upstream_exception_w return mock_response mock_logging_obj = _make_logging_obj() - provider_config = MagicMock() received = [] async def _drain(): async for chunk in AsyncPassthroughStreamingResponse( response=response_coro(), litellm_logging_obj=mock_logging_obj, - provider_config=provider_config, + provider_config=_provider_config(_RecordingCollector()), ): received.append(chunk) @@ -196,11 +212,7 @@ async def test_asyncpassthroughstreamingresponse_flushes_on_upstream_exception_w await asyncio.sleep(0) - mock_logging_obj.async_flush_passthrough_collected_chunks.assert_called_once() - call_kwargs = ( - mock_logging_obj.async_flush_passthrough_collected_chunks.call_args.kwargs - ) - assert call_kwargs["raw_bytes"] == partial_chunks + assert _spend_payload(mock_logging_obj.async_flush_passthrough_collected_chunks) == b"".join(partial_chunks) def test_passthroughstreamingresponse_flushes_on_normal_completion(): @@ -221,12 +233,11 @@ def test_passthroughstreamingresponse_flushes_on_normal_completion(): mock_logging_obj = MagicMock() mock_logging_obj.flush_passthrough_collected_chunks = MagicMock() - provider_config = MagicMock() received_responce = PassthroughStreamingResponse( response=mock_response, litellm_logging_obj=mock_logging_obj, - provider_config=provider_config, + provider_config=_provider_config(_RecordingCollector()), ) with patch("litellm.utils.executor", _ImmediateExecutor()): @@ -237,7 +248,7 @@ def test_passthroughstreamingresponse_flushes_on_normal_completion(): assert received_responce.headers["content-type"] == "application/octet-stream" assert received_responce.headers["x-request-id"] == "req-123" - mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once() + assert _spend_payload(mock_logging_obj.flush_passthrough_collected_chunks) == b"".join(chunks) def test_passthroughstreamingresponse_flushes_on_early_close(): @@ -258,19 +269,66 @@ def test_passthroughstreamingresponse_flushes_on_early_close(): mock_logging_obj = MagicMock() mock_logging_obj.flush_passthrough_collected_chunks = MagicMock() - provider_config = MagicMock() with patch("litellm.utils.executor", _ImmediateExecutor()): gen = PassthroughStreamingResponse( response=mock_response, litellm_logging_obj=mock_logging_obj, - provider_config=provider_config, + provider_config=_provider_config(_RecordingCollector()), ) first = next(gen) gen.close() assert first == chunks[0] - mock_logging_obj.flush_passthrough_collected_chunks.assert_called_once() - call_kwargs = mock_logging_obj.flush_passthrough_collected_chunks.call_args.kwargs - assert call_kwargs["raw_bytes"] == [chunks[0]] + assert _spend_payload(mock_logging_obj.flush_passthrough_collected_chunks) == chunks[0] + + +@pytest.mark.asyncio +async def test_asyncpassthroughstreamingresponse_relays_the_stream_when_spend_parsing_fails(): + from litellm.passthrough.main import AsyncPassthroughStreamingResponse + + chunks = [b"chunk-1", b"chunk-2", b"chunk-3"] + mock_response = _make_streaming_response(chunks) + + async def response_coro(): + return mock_response + + mock_logging_obj = _make_logging_obj() + + received = [ + chunk + async for chunk in AsyncPassthroughStreamingResponse( + response=response_coro(), + litellm_logging_obj=mock_logging_obj, + provider_config=_provider_config(_FailingCollector()), + ) + ] + await asyncio.sleep(0) + + assert received == chunks + mock_logging_obj.async_flush_passthrough_collected_chunks.assert_not_called() + + +def test_passthroughstreamingresponse_relays_the_stream_when_spend_parsing_fails(): + from litellm.passthrough.main import PassthroughStreamingResponse + + chunks = [b"a", b"b", b"c"] + mock_response = MagicMock(spec=httpx.Response) + mock_response.status_code = 200 + mock_response.headers = httpx.Headers({"content-type": "application/octet-stream"}) + mock_response.iter_bytes = lambda: iter(chunks) + + mock_logging_obj = MagicMock() + mock_logging_obj.flush_passthrough_collected_chunks = MagicMock() + + received = list( + PassthroughStreamingResponse( + response=mock_response, + litellm_logging_obj=mock_logging_obj, + provider_config=_provider_config(_FailingCollector()), + ) + ) + + assert received == chunks + mock_logging_obj.flush_passthrough_collected_chunks.assert_not_called() diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_classify.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_classify.py index dc20d664a53..30107db4055 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_classify.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_classify.py @@ -1,6 +1,9 @@ """Classification matrix for upstream OAuth/DCR rejections: who is blamed depends only on the §5.2 code and whose credentials the gateway presented, never on the upstream's HTTP status.""" +from typing import Final + +import pytest import httpx from litellm.proxy._experimental.mcp_server.faults.classify import ( @@ -12,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import ( GatewayRejected, UpstreamProtocolFault, UpstreamReportedFault, + UpstreamRegistrationRefused, ) @@ -145,3 +149,28 @@ def test_dcr_server_error_code_is_not_blamed_on_caller(): log_context="srv", ) assert isinstance(fault, UpstreamReportedFault) + + +@pytest.mark.parametrize("status_code", [401, 403]) +@pytest.mark.parametrize("body", ["Forbidden", 'private upstream details', '{"error": ""}', '{"error": 12}']) +def test_dcr_access_refusal_without_oauth_error(status_code: int, body: str) -> None: + fault: Final = classify_upstream_dcr_rejection(_response(status_code, text_body=body), log_context="srv") + assert isinstance(fault, UpstreamRegistrationRefused) + assert fault.status_code == status_code + + +@pytest.mark.parametrize("status_code", [401, 403]) +def test_dcr_access_refusal_preserves_oauth_error(status_code: int) -> None: + fault: Final = classify_upstream_dcr_rejection( + _response(status_code, json_body={"error": "invalid_redirect_uri", "error_description": "not allowed"}), + log_context="srv", + ) + assert fault == CallerRejected(code="invalid_redirect_uri", description="not allowed") + + +@pytest.mark.parametrize("status_code", [401, 403]) +def test_token_access_refusal_remains_protocol_fault(status_code: int) -> None: + fault: Final = classify_upstream_token_rejection( + _response(status_code, text_body="Forbidden"), credential_source="gateway_stored", log_context="srv" + ) + assert isinstance(fault, UpstreamProtocolFault) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_render_oauth.py b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_render_oauth.py index 78513e315a7..a6807ae1454 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_render_oauth.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/faults/test_render_oauth.py @@ -2,6 +2,9 @@ code can never ship on a server-fault status and gateway-side faults never carry provider prose.""" import json +from typing import Final, Literal + +import pytest from litellm.proxy._experimental.mcp_server.faults.render_oauth import ( dcr_fault_detail, @@ -12,6 +15,7 @@ from litellm.proxy._experimental.mcp_server.faults.types import ( GatewayRejected, UpstreamProtocolFault, UpstreamReportedFault, + UpstreamRegistrationRefused, ) @@ -94,3 +98,17 @@ def test_dcr_upstream_reported_fault_maps_to_5xx(): status_code, detail = dcr_fault_detail(UpstreamReportedFault(code="server_error")) assert status_code == 502 assert "internal error" in detail + + +@pytest.mark.parametrize("upstream_status", [401, 403]) +def test_registration_refusal_gives_configuration_guidance(upstream_status: Literal[401, 403]) -> None: + fault: Final = UpstreamRegistrationRefused(status_code=upstream_status) + status, detail = dcr_fault_detail(fault) + assert status == 403 + assert f"HTTP {upstream_status}" in detail + assert "may require a pre-registered OAuth client" in detail + assert "client_id" in detail and "client_secret" in detail + response: Final = render_token_fault(fault) + assert response.status_code == 400 + assert json.loads(response.body) == {"error": "unauthorized_client", "error_description": detail} + assert response.headers["cache-control"] == "no-store" diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py index 010e7e14d39..774cd022703 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py @@ -478,3 +478,47 @@ async def test_bearer_auth_advertises_the_header_it_will_occupy(): assert ClientCredentialsBearerAuth("t", refetch, ClientCredentialsConfig()).header_name == "Authorization" default_carrier = ClientCredentialsConfig(header_name="esb-oauth") assert ClientCredentialsBearerAuth("t", refetch, default_carrier).header_name == "esb-oauth" + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["denied", "invalid", "missing", "success", "timeout", "connect", "cancel"]) +async def test_token_exchange_failure_diagnostics(mode, monkeypatch, caplog): + import asyncio + import logging + from litellm.llms.custom_httpx import http_handler + from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import post_client_credentials_grant + + class Poster: + async def post(self, url, headers, data): + request = httpx.Request("POST", url, headers=headers, data=data) + if mode == "timeout": + raise httpx.ReadTimeout("private-transport-message", request=request) + if mode == "connect": + raise httpx.ConnectError("private-transport-message", request=request) + if mode == "cancel": + raise asyncio.CancelledError + response = httpx.Response(401 if mode == "denied" else 200, request=request, + content=b"not-json-private" if mode == "invalid" else None, + json=None if mode == "invalid" else {"error": "invalid_client", "client_secret":"first second", **({"access_token":"private-token"} if mode == "success" else {})}) + response.raise_for_status() + return response + + monkeypatch.setattr(http_handler, "get_async_httpx_client", lambda **kwargs: Poster()) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + if mode == "cancel": + with pytest.raises(asyncio.CancelledError): + await post_client_credentials_grant("https://idp/token", {}, {}) + assert not caplog.text + return + result = await post_client_credentials_grant("https://idp/token?key=query-secret", {"client_secret":"first second"}, {"X-Custom":"header-secret"}) + for secret in ("first", "second", "query-secret", "header-secret", "private-token", "not-json-private", "private-transport-message"): + assert secret not in caplog.text + if mode == "success": + assert isinstance(result, TokenEndpointSuccess) and result.body["access_token"] == "private-token" + assert not caplog.text + elif mode in {"timeout", "connect"}: + assert isinstance(result, TokenEndpointUnreachable) + assert "POST https://idp/ failed" in caplog.text + else: + assert "POST https://idp/ -> HTTP" in caplog.text + assert {"denied":"denied", "invalid":"invalid response", "missing":"no access token"}[mode] in caplog.text diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py index 5a99139a67f..1b7e3d6a1c3 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_discoverable_endpoints.py @@ -7150,7 +7150,7 @@ async def test_extract_user_id_rehydrates_cross_replica_dict_cache(proxy_globals key = "sk-alice-key" cache = UserApiKeyCache() - cache.in_memory_cache.set_cache(hash_token(key), {"token": hash_token(key), "user_id": "alice"}) + cache.set_cache(hash_token(key), {"token": hash_token(key), "user_id": "alice"}) proxy_globals.user_api_key_cache = cache proxy_globals.prisma_client = object() @@ -11141,3 +11141,48 @@ async def test_enforced_login_warms_verified_token_readable_without_database_loo assert token.identity_binding_proof == proof assert token.refresh_token is None read.assert_not_awaited() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("upstream_status", [401, 403]) +@pytest.mark.parametrize("auth_type", [MCPAuth.true_passthrough, MCPAuth.oauth_delegate]) +@pytest.mark.parametrize("dcr_bridge", [False, True]) +@pytest.mark.parametrize("flow", ["register", "mint"]) +async def test_dcr_refusal_is_actionable_without_upstream_body( + upstream_status: int, auth_type: MCPAuth, dcr_bridge: bool, flow: str, monkeypatch: pytest.MonkeyPatch +) -> None: + import httpx + from typing import Final + + from litellm.proxy._experimental.mcp_server.discoverable_endpoints import ( + mint_ephemeral_dcr_client, + register_client_with_server, + ) + + server: Final = _bridge_server( + auth_type=auth_type, dcr_bridge=dcr_bridge, server_id=f"refused-{auth_type}-{dcr_bridge}-{flow}-{upstream_status}", + client_id=None, + ) + import respx + + monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") + with respx.mock as upstream: + registration: Final = upstream.post(server.registration_url).mock( + return_value=httpx.Response(upstream_status, text="Forbidden private upstream details") + ) + operation: Final = ( + mint_ephemeral_dcr_client(_bridge_mock_request(), server) + if flow == "mint" + else register_client_with_server( + request=_bridge_mock_request(), mcp_server=server, client_name="Test client", + grant_types=None, response_types=None, token_endpoint_auth_method=None, + client_redirect_uris=["http://localhost:9999/callback"], + ) + ) + with pytest.raises(HTTPException) as exc: + await operation + assert registration.call_count == 1 + assert exc.value.status_code == 403 + assert f"HTTP {upstream_status}" in str(exc.value.detail) + assert "pre-registered OAuth client" in str(exc.value.detail) + assert "private upstream details" not in str(exc.value.detail) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py index 7d46bb0237a..b6535e6326a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_debug.py @@ -10,9 +10,13 @@ from starlette.types import Message from litellm.proxy._experimental.mcp_server.outbound_credentials.types import AuthResolution +import httpx + from litellm.proxy._experimental.mcp_server.mcp_debug import ( MCP_DEBUG_REQUEST_HEADER, MCPDebug, + describe_upstream_http_failure, + MCPAuthDiagnostics, ) @@ -206,6 +210,206 @@ class TestWrapSendWithDebugHeaders: assert captured[0] == body_msg +class TestDescribeUpstreamHttpFailure: + @staticmethod + def _status_error(*, body: bytes, response_body: bytes | None = None) -> httpx.HTTPStatusError: + request = httpx.Request( + "POST", + "https://upstream.example/apis/mcp", + headers={"Authorization": "Bearer secret-token-abcdef0123456789", "Content-Type": "application/json" if body.startswith(b"{") else "application/x-www-form-urlencoded"}, + content=body, + ) + response = ( + httpx.Response(500, request=request, content=response_body) + if response_body is not None + else httpx.Response(500, request=request, stream=httpx.ByteStream(b'{"error":"boom"}')) + ) + return httpx.HTTPStatusError("500", request=request, response=response) + + def test_includes_method_url_status_and_request_body(self): + exc = self._status_error( + body=b'{"method":"initialize","jsonrpc":"2.0","id":0}', + response_body=b'{"error":"boom"}', + ) + described = describe_upstream_http_failure(exc) + assert described is not None + assert "POST https://upstream.example/ -> HTTP 500" in described + assert '{"method":"initialize"' in described + assert 'response body: {"error":"boom"}' in described + + def test_masks_authorization_header_and_secret_body_fields(self): + exc = self._status_error( + body=b"grant_type=client_credentials&client_id=abc&client_secret=super-secret-value-1234", + response_body=b"{}", + ) + described = describe_upstream_http_failure(exc) + assert described is not None + assert "secret-token-abcdef0123456789" not in described + assert "super-secret-value-1234" not in described + assert "client_id=abc" in described + assert "client_secret=" in described + + def test_reports_unread_streamed_response_body(self): + described = describe_upstream_http_failure(self._status_error(body=b"{}")) + assert described is not None + assert "response body: (not read)" in described + + def test_finds_response_behind_cause_chain(self): + wrapper = RuntimeError("token minting failed") + wrapper.__cause__ = self._status_error(body=b"{}", response_body=b'{"error":"invalid_client"}') + described = describe_upstream_http_failure(wrapper) + assert described is not None + assert "invalid_client" in described + + def test_returns_none_without_http_response(self): + assert describe_upstream_http_failure(ConnectionError("refused")) is None + + +@pytest.mark.parametrize("body", [ + b'{"password":"first second","token":"demo-secret"}', + b'{"nested":[{"access_token":"first,second"}]}', + b'client%5Fsecret=first+second&token=demo-secret', +]) +def test_failure_log_fully_redacts_structured_secrets(body): + request = httpx.Request("POST", "https://upstream/mcp?credential=query-secret", + headers={"X-Custom-Credential": "custom-secret"}, content=body) + response = httpx.Response(500, request=request, content=body) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response)) + assert detail is not None + for secret in ("first", "second", "demo-secret", "custom-secret", "query-secret"): + assert secret not in detail + + +def test_failure_log_omits_unstructured_body(): + request = httpx.Request("POST", "https://upstream/mcp", content=b"arbitrary-secret") + response = httpx.Response(500, request=request, content=b"arbitrary-secret") + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response)) + assert detail is not None + assert "arbitrary-secret" not in detail + assert "omitted" in detail + + +@pytest.mark.asyncio +@pytest.mark.parametrize("mode", ["error", "empty", "large", "timeout", "read_failure", "closed", "success", "cancel"]) +async def test_error_capture_is_bounded_and_preserves_success_and_cancellation(mode): + from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response + + class Stream(httpx.AsyncByteStream): + def __init__(self): + self.reads = 0 + + async def __aiter__(self): + self.reads += 1 + if mode == "timeout": + await asyncio.sleep(10) + if mode == "closed": + raise httpx.StreamClosed() + if mode == "read_failure": + raise httpx.ReadError("private-read-error") + if mode == "cancel": + raise asyncio.CancelledError + yield b"" if mode == "empty" else b'{"error":"missing_scope","password":"first second"}' if mode != "large" else b"x" * 20000 + + stream = Stream() + request = httpx.Request("POST", "https://upstream/mcp") + response = httpx.Response(200 if mode == "success" else 500, request=request, stream=stream) + if mode == "cancel": + with pytest.raises(asyncio.CancelledError): + await capture_upstream_error_response(response) + return + await capture_upstream_error_response(response) + if mode == "success": + assert stream.reads == 0 + assert await response.aread() == b'{"error":"missing_scope","password":"first second"}' + return + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response)) + assert detail is not None + assert "first" not in detail and "second" not in detail and "private-read-error" not in detail + expected = {"empty": "(empty)", "error": "missing_scope", "large": "capture limit", "timeout": "read failed", "read_failure": "read failed", "closed":"read failed"} + assert expected[mode] in detail + if mode == "error": + assert await response.aread() == b'{"error":"missing_scope","password":"first second"}' + + +@pytest.mark.parametrize("body", [b"", b'"scalar"', b'{"hint":"line1\\nline2"}', b'{"hint":"' + b'x' * 600 + b'"}']) +def test_failure_preview_handles_empty_scalar_control_and_long_bodies(body): + request = httpx.Request("POST", "https://user:secret@upstream/mcp?key=private#private", content=body) + response = httpx.Response(500, request=request, content=body) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response)) + assert detail is not None + assert "private" not in detail and "user:secret" not in detail and "\n" not in detail + if not body: + assert "(empty)" in detail + elif body.startswith(b'"'): + assert "omitted" in detail + elif len(body) > 512: + assert "truncated" in detail and len(detail) < 1300 + else: + assert "line1\\nline2" in detail + + +@pytest.mark.asyncio +@pytest.mark.parametrize("slow_error", [False, True]) +async def test_error_capture_preserves_httpx_auth_retry(slow_error): + from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response + + class RetryAuth(httpx.Auth): + def auth_flow(self, request): + response = yield request + if response.status_code == 401: + request.headers["Authorization"] = "Bearer refreshed" + yield request + + class SlowStream(httpx.AsyncByteStream): + async def __aiter__(self): + await asyncio.sleep(10) + yield b'{"error":"expired_token"}' + + def upstream(request): + if request.headers.get("Authorization"): + return httpx.Response(200, json={"ok": True}) + return httpx.Response(401, stream=SlowStream()) if slow_error else httpx.Response(401, json={"error":"expired_token"}) + + async with httpx.AsyncClient(transport=httpx.MockTransport(upstream), auth=RetryAuth(), + event_hooks={"response":[capture_upstream_error_response]}) as client: + response = await client.get("https://upstream/mcp") + assert response.status_code == 200 and response.json() == {"ok":True} + if slow_error: + assert response.history[0].content == b"" + else: + assert response.history[0].json() == {"error":"expired_token"} + + +def test_failure_diagnostics_without_request_and_with_streamed_request(): + response = httpx.Response(503) + exc = httpx.HTTPStatusError("failed", request=httpx.Request("GET", "https://upstream"), response=response) + assert describe_upstream_http_failure(exc) == "HTTP 503 | request unavailable" + request = httpx.Request("POST", "https://upstream", content=iter((b"private-body",))) + response = httpx.Response(503, request=request) + described = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response)) + assert described is not None and "streamed, not captured" in described and "private-body" not in described + + + +def test_deep_error_body_is_bounded_without_exposing_nested_values(): + body = b'{"nested":' * 18 + b'{"password":"hidden-value"}' + b'}' * 18 + request = httpx.Request("POST", "https://upstream/mcp", content=body) + response = httpx.Response(500, request=request, content=body) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response)) + assert detail is not None and "hidden-value" not in detail + assert "nested" in detail and "REDACTED" in detail + + +@pytest.mark.parametrize("body", [b'client%5Fsecret=first+second&client_id=visible', b'client_secret=first%26second&client_id=visible']) +def test_encoded_form_credentials_are_decoded_before_redaction(body): + request = httpx.Request("POST", "https://upstream/token", content=body, + headers={"Content-Type":"application/x-www-form-urlencoded"}) + response = httpx.Response(400, request=request, content=body, + headers={"Content-Type":"application/x-www-form-urlencoded"}) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failure", request=request, response=response)) + assert detail is not None and "client_id=visible" in detail + assert "first" not in detail and "second" not in detail + @pytest.mark.asyncio @pytest.mark.parametrize("source", tuple(AuthResolution)) @pytest.mark.parametrize("method", ("GET", "DELETE", "POST")) @@ -291,3 +495,99 @@ async def test_concurrent_mcp_messages_record_on_their_own_http_scope() -> None: await asyncio.gather(record(first, AuthResolution.stored_user_token), record(second, AuthResolution.per_request_header)) assert first.resolution() == "stored-user-token" assert second.resolution() == "per-request-header" + + +@pytest.mark.parametrize("source", ["header", "bearer", "basic", "cookie", "query", "form", "json"]) +def test_reflected_credentials_are_removed_from_normal_response_fields(source): + import base64 + + secret = "generic-credential-123" + headers = {"X-Custom":secret} if source == "header" else {"Authorization":"Bearer " + secret} if source == "bearer" else {"Authorization":"Basic " + base64.b64encode(("client:" + secret).encode()).decode()} if source == "basic" else {"Cookie":"session=" + secret} if source == "cookie" else {} + request = httpx.Request("POST", "https://upstream/token" + ("?credential=" + secret if source == "query" else ""), + headers=headers, data={"client_secret":secret} if source == "form" else None, + json={"nested":{"client_secret":secret}} if source == "json" else None) + response = httpx.Response(401, request=request, json={"error":"invalid_client", "error_description":"Rejected " + secret}) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response)) + assert detail is not None and "invalid_client" in detail + assert secret not in detail and "REDACTED" in detail + + + +@pytest.mark.parametrize("secret", ['value"with\ncharacters€', "R"]) +def test_reflected_values_are_redacted_before_truncation_without_expanding_replacements(secret): + request = httpx.Request("POST", "https://upstream/token", json={"client_secret":secret}) + response = httpx.Response(401, request=request, json={"error":"invalid_client", "detail":"x" * 460 + secret}) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response)) + assert detail is not None and "invalid_client" in detail + assert "value" not in detail and "characters" not in detail and len(detail) < 1400 + + +@pytest.mark.parametrize("headers", [{"Authorization":"Basic !!!"}, {"Cookie":"bad@key=opaque"}]) +def test_malformed_auth_headers_do_not_break_failure_diagnostics(headers): + request = httpx.Request("POST", "https://upstream/token", headers=headers) + response = httpx.Response(401, request=request, json={"error":"invalid_client"}) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response)) + assert detail is not None and "invalid_client" in detail + assert "!!!" not in detail and "opaque" not in detail + + +def test_oversized_request_omits_potentially_reflected_response_credentials(): + request = httpx.Request("POST", "https://upstream/token", content=b"x" * 17000) + response = httpx.Response(401, request=request, json={"error_description":"unknown-secret"}) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response)) + assert detail is not None and "capture limit" in detail and "credentials unavailable" in detail + assert "unknown-secret" not in detail + + + +@pytest.mark.asyncio +async def test_streamed_error_redacts_reflected_credentials_before_capture(): + import json + from litellm.proxy._experimental.mcp_server.mcp_debug import capture_upstream_error_response + + secret = "generic-credential-123" + request = httpx.Request("POST", "https://upstream/token", data={"client_secret":secret}) + raw = json.dumps({"error":"invalid_client", "error_description":"Rejected " + secret}).encode() + response = httpx.Response(401, request=request, stream=httpx.ByteStream(raw)) + await capture_upstream_error_response(response) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response)) + assert detail is not None and "invalid_client" in detail and "Rejected" in detail + assert secret not in detail and "REDACTED" in detail + assert await response.aread() == raw + + +@pytest.mark.parametrize("path", ["/credential-path-value/mcp", "/oauth/credential-path-value/token"]) +def test_failure_diagnostics_omit_credential_bearing_url_paths(path): + request = httpx.Request("POST", "https://upstream.example" + path) + response = httpx.Response(401, request=request, json={"error": "access_denied"}) + error = httpx.HTTPStatusError("denied", request=request, response=response) + diagnostic = describe_upstream_http_failure(error) + assert diagnostic is not None + assert "credential-path-value" not in diagnostic + assert "POST https://upstream.example/ -> HTTP 401" in diagnostic + assert "access_denied" in diagnostic + + +def test_deep_request_omits_response_when_credentials_cannot_be_inspected(): + from litellm.proxy._experimental.mcp_server.utils import MAX_STRUCTURED_CONTENT_SCAN_DEPTH + + raw = "[" * (MAX_STRUCTURED_CONTENT_SCAN_DEPTH + 1) + '{"client_secret":"nested-credential"}' + "]" * (MAX_STRUCTURED_CONTENT_SCAN_DEPTH + 1) + request = httpx.Request("POST", "https://upstream/token", content=raw, headers={"Content-Type": "application/json"}) + response = httpx.Response(401, request=request, json={"error_description": "Rejected nested-credential"}) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response)) + assert detail is not None and "HTTP 401" in detail + assert "response body: (omitted: request credentials unavailable)" in detail + assert "nested-credential" not in detail + + +@pytest.mark.parametrize("field", ["accessToken", "refreshToken", "clientSecret", "apikey", "CLIENTASSERTION", "cost_token"]) +@pytest.mark.parametrize("encoding", ["json", "form"]) +def test_compact_credential_fields_and_reflected_values_are_redacted(field, encoding): + secret = "generic-private-value" + fields = {field: secret} + request = httpx.Request("POST", "https://upstream/token", json=fields if encoding == "json" else None, + data=fields if encoding == "form" else None) + response = httpx.Response(401, request=request, json={field: secret, "error": "invalid_client", "detail": "Rejected " + secret}) + detail = describe_upstream_http_failure(httpx.HTTPStatusError("failed", request=request, response=response)) + assert detail is not None and "invalid_client" in detail + assert "REDACTED" in detail and secret not in detail diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py index 9667224de98..3f5d4ad83ea 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_oauth_passthrough_tools.py @@ -1,5 +1,6 @@ """Unit tests for MCP OAuth passthrough tool-fetch behavior.""" +import logging import sys from unittest.mock import AsyncMock, MagicMock @@ -11,7 +12,7 @@ if sys.version_info < (3, 11): from exceptiongroup import ExceptionGroup -from litellm.proxy._experimental.mcp_server.exceptions import MCPUpstreamAuthError +from litellm.proxy._experimental.mcp_server.exceptions import MCPServerListError, MCPUpstreamAuthError from litellm.proxy._experimental.mcp_server.mcp_server_manager import ( MCPServerManager, _extract_upstream_auth_failure, @@ -434,3 +435,43 @@ async def test_aggregate_with_single_accessible_server_still_absorbs(): assert listing.tools == [] assert listing.outcomes["delegate_docs"].tag == "auth_required" + + +@pytest.mark.asyncio +async def test_fetch_tools_logs_upstream_request_details_on_500(caplog): + manager = MCPServerManager() + request = httpx.Request( + "POST", + "https://upstream/apis/mcp", + headers={"Authorization": "Bearer upstream-token-0123456789"}, + content=b'{"method":"initialize","jsonrpc":"2.0","id":0}', + ) + response = httpx.Response(500, request=request) + mock_client = MagicMock() + mock_client.list_tools = AsyncMock( + side_effect=httpx.HTTPStatusError("500", request=request, response=response) + ) + + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + with pytest.raises(MCPServerListError): + await manager._fetch_tools_with_timeout(mock_client, "sample_docs") + + assert "POST https://upstream/ -> HTTP 500" in caplog.text + assert '"method":"initialize"' in caplog.text + assert "upstream-token-0123456789" not in caplog.text + + + +@pytest.mark.asyncio +async def test_client_creation_failure_logs_sanitized_exchange(monkeypatch, caplog): + manager = MCPServerManager() + server = MCPServer(server_id="sample", name="sample", url="https://upstream/mcp", transport=MCPTransport.http, auth_type=MCPAuth.none) + request = httpx.Request("POST", "https://upstream/mcp?credential=query-secret") + response = httpx.Response(500, request=request, json={"error":"missing_scope"}) + error = httpx.HTTPStatusError("query-secret", request=request, response=response) + monkeypatch.setattr(manager, "_create_mcp_client", AsyncMock(side_effect=error)) + with caplog.at_level(logging.WARNING, logger="LiteLLM"): + with pytest.raises(MCPServerListError): + await manager._get_tools_from_server(server) + assert "POST https://upstream/ -> HTTP 500" in caplog.text + assert "missing_scope" in caplog.text and "query-secret" not in caplog.text diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py index a6e43686e2c..b637586d7ff 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_mcp_server_manager.py @@ -12787,6 +12787,125 @@ async def test_debug_reports_legacy_signing_and_non_http_transport(transport: Li finally: request_ctx.reset(token) + +@pytest.mark.asyncio +async def test_temporary_server_discovery_reuses_resolved_metadata_without_publishing() -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="temporary-oauth-discovery", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + ) + manager._set_oauth_discovery_deferred(server.server_id, True) + metadata: Final = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + registration_url="https://idp.example.com/register", + ) + with patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery: + resolved: Final = await manager.ensure_oauth_metadata_discovered(server) + repeated: Final = await manager.ensure_oauth_metadata_discovered(server) + assert resolved.authorization_url == metadata.authorization_url + assert resolved.token_url == metadata.token_url + assert resolved.registration_url == metadata.registration_url + assert repeated is resolved + assert server.server_id not in manager.registry + assert server.server_id not in manager.config_mcp_servers + discovery.assert_awaited_once() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("auth_type", [MCPAuth.oauth2, MCPAuth.true_passthrough]) +async def test_repeated_stale_oauth_discovery_is_bounded(auth_type: MCPAuth) -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="repeated-stale", name="stale", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=auth_type, oauth2_flow="authorization_code", + ) + manager.registry[server.server_id] = server + manager._set_oauth_discovery_deferred(server.server_id, True) + metadata: Final = MCPOAuthMetadata( + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + ) + with ( + patch.object(manager, "_discover_oauth_metadata_for_server", AsyncMock(return_value=metadata)) as discovery, + patch.object(manager, "_publish_resolved_oauth_server", return_value=None), + ): + if auth_type == MCPAuth.true_passthrough: + assert await manager.ensure_oauth_metadata_discovered(server) is server + else: + with pytest.raises(HTTPException) as exc: + await manager.ensure_oauth_metadata_discovered(server) + assert exc.value.status_code == 503 + assert "changed repeatedly" in str(exc.value.detail) + assert discovery.await_count == 2 + + +@pytest.mark.asyncio +async def test_stale_discovery_falls_back_to_resolved_registered_server() -> None: + manager: Final = MCPServerManager() + original: Final = MCPServer( + server_id="resolved-replacement", name="replacement", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, oauth2_flow="authorization_code", + ) + replacement: Final = original.model_copy(update={ + "url": "https://new.example.com/mcp", "authorization_url": "https://new.example.com/authorize", + "token_url": "https://new.example.com/token", + }) + manager.registry[original.server_id] = replacement + assert await manager._rejoin_oauth_metadata_discovery(original, retry_stale=False) is replacement + + +def test_stale_discovery_cannot_overwrite_new_registered_server() -> None: + manager: Final = MCPServerManager() + original: Final = MCPServer( + server_id="stale-publication", name="publication", url="https://old.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.oauth2, + ) + manager._set_oauth_discovery_deferred(original.server_id, True) + original_slot: Final = manager._oauth_discovery_slot(original.server_id) + assert original_slot is not None + replacement: Final = original.model_copy(update={"url": "https://new.example.com/mcp"}) + manager.registry[original.server_id] = replacement + manager._set_oauth_discovery_deferred(original.server_id, True) + assert manager._publish_resolved_oauth_server(original, original_slot.generation) is None + assert manager.registry[original.server_id] is replacement + + +@pytest.mark.asyncio +async def test_temporary_oauth_discovery_expires_without_more_requests() -> None: + manager: Final = MCPServerManager() + server: Final = MCPServer( + server_id="expiring-session", name="temporary", url="https://idp.example.com/mcp", + transport=MCPTransport.http, auth_type=MCPAuth.true_passthrough, + authorization_url="https://idp.example.com/authorize", token_url="https://idp.example.com/token", + ) + manager._set_oauth_discovery_deferred(server.server_id, True) + resolved: Final = await manager.ensure_oauth_metadata_discovered(server) + assert manager._oauth_discovery_slot(server.server_id) is not None + loop: Final = asyncio.get_running_loop() + expired: Final = loop.create_future() + with patch.object(loop, "time", return_value=loop.time() + 301): + loop.call_later(0, expired.set_result, None) + await expired + assert resolved.authorization_url == server.authorization_url + assert manager._oauth_discovery_slot(server.server_id) is None + + +def test_old_temporary_discovery_expiry_preserves_replacement() -> None: + manager: Final = MCPServerManager() + manager._set_oauth_discovery_deferred("reused-session", True) + old_slot: Final = manager._oauth_discovery_slot("reused-session") + assert old_slot is not None + manager._set_oauth_discovery_deferred("reused-session", True) + replacement: Final = manager._oauth_discovery_slot("reused-session") + manager._expire_temporary_oauth_discovery("reused-session", old_slot.generation) + assert manager._oauth_discovery_slot("reused-session") is replacement + assert replacement is not None + manager._expire_temporary_oauth_discovery("reused-session", replacement.generation) + assert manager._oauth_discovery_slot("reused-session") is None + manager._expire_temporary_oauth_discovery("reused-session", replacement.generation) + assert manager._oauth_discovery_slot("reused-session") is None + + @pytest.mark.asyncio async def test_openapi_health_coalesces_concurrent_checks_and_reuses_results(respx_mock, monkeypatch): monkeypatch.setenv("DISABLE_AIOHTTP_TRANSPORT", "True") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py index 16c6aa128d0..31ccd5c9817 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/test_rest_endpoints.py @@ -2595,6 +2595,79 @@ class TestCallToolRestAPI: assert result == masked_result + async def test_success_logging_start_time_excludes_pre_call_processing(self, monkeypatch): + """Pre-call hook latency (guardrails, header resolution) must not inflate the tool call's + logged duration on success.""" + from litellm.proxy import proxy_server + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + server_id = "server-1" + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + auth_type = None + + stub_server = StubServer() + + async def fake_add_litellm_data_to_request(**kwargs): + return kwargs.get("data", {}) + + pre_call_finished_at = {} + + async def slow_pre_call_hook(user_api_key_dict, data, call_type): + await asyncio.sleep(0.05) + pre_call_finished_at["value"] = datetime.now() + return data + + captured = {} + + async def fake_execute_mcp_tool(**kwargs): + captured.update(kwargs) + return {"result": "ok"} + + fire_logging = AsyncMock(return_value={"result": "ok"}) + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: stub_server if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + proxy_server, "add_litellm_data_to_request", fake_add_litellm_data_to_request, raising=False + ) + monkeypatch.setattr(proxy_server, "proxy_config", {}, raising=False) + monkeypatch.setattr(proxy_server.proxy_logging_obj, "pre_call_hook", slow_pre_call_hook) + monkeypatch.setattr(rest_endpoints, "execute_mcp_tool", fake_execute_mcp_tool, raising=False) + monkeypatch.setattr(rest_endpoints, "_fire_mcp_tool_call_logging", fire_logging, raising=False) + + request = _build_request( + path="/mcp-rest/tools/call", + method="POST", + json_body={"server_id": "server-1", "name": "demo-tool", "arguments": {}}, + ) + + await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=UserAPIKeyAuth()) + + logged_start_time = fire_logging.await_args.args[2] + assert captured["start_time"] >= pre_call_finished_at["value"] + assert logged_start_time == captured["start_time"] + async def test_success_logging_guardrail_rejection_propagates(self, monkeypatch): """A guardrail rejecting the tool result must not be swallowed as a logging failure, otherwise the unguarded result would still be returned to the caller.""" @@ -2765,6 +2838,139 @@ class TestCallToolRestAPI: info_messages = [_rendered_log_message(c) for c in mock_logger.info.call_args_list if c.args] assert not any("relaying upstream" in m for m in info_messages) + @pytest.mark.parametrize("raise_site", ["pre_call_hook", "execute_mcp_tool"]) + async def test_guardrail_block_runs_failure_logging_before_http_translation(self, monkeypatch, raise_site): + """A pre_mcp_call guardrail block, whether raised by the pre-call hook or from inside + execute_mcp_tool, must reach proxy_logging_obj.post_call_failure_hook (the only path that + writes the failure spend-log row) with the logging object's failure payload already built, + and the REST caller must still get the same 400 it got before.""" + from litellm.proxy import proxy_server + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_get_allowed_mcp_servers(*args, **kwargs): + return ["server-1"] + + class StubServer: + server_id = "server-1" + alias = "server-1" + server_name = "server-1" + name = "stub" + allowed_tools = None + mcp_info = {"server_name": "stub"} + available_on_public_internet = True + auth_type = None + + async def fake_add_litellm_data_to_request(**kwargs): + return kwargs.get("data", {}) + + guardrail_error = HTTPException( + status_code=400, + detail={"error": "Content blocked: keyword 'confidential' detected", "keyword": "confidential"}, + ) + + async def passthrough_pre_call_hook(user_api_key_dict, data, call_type): + return data + + async def blocking_pre_call_hook(user_api_key_dict, data, call_type): + raise guardrail_error + + async def fake_execute_mcp_tool(**kwargs): + raise guardrail_error + + async def passthrough_execute_mcp_tool(**kwargs): + return [] + + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_allowed_mcp_servers", + fake_get_allowed_mcp_servers, + raising=False, + ) + monkeypatch.setattr( + rest_endpoints.global_mcp_server_manager, + "get_mcp_server_by_id", + lambda server_id: StubServer() if server_id == "server-1" else None, + raising=False, + ) + monkeypatch.setattr( + proxy_server, "add_litellm_data_to_request", fake_add_litellm_data_to_request, raising=False + ) + monkeypatch.setattr(proxy_server, "proxy_config", {}, raising=False) + monkeypatch.setattr( + proxy_server.proxy_logging_obj, + "pre_call_hook", + blocking_pre_call_hook if raise_site == "pre_call_hook" else passthrough_pre_call_hook, + ) + monkeypatch.setattr( + rest_endpoints, + "execute_mcp_tool", + fake_execute_mcp_tool if raise_site == "execute_mcp_tool" else passthrough_execute_mcp_tool, + raising=False, + ) + post_call_failure_hook = AsyncMock(return_value=None) + monkeypatch.setattr(proxy_server.proxy_logging_obj, "post_call_failure_hook", post_call_failure_hook) + + user_api_key_dict = UserAPIKeyAuth(api_key="hashed-key", request_route="/mcp-rest/tools/call") + request = _build_request( + headers={"x-mcp-deepwiki-authorization": "Bearer upstream-secret"}, + path="/mcp-rest/tools/call", + method="POST", + json_body={"server_id": "server-1", "name": "demo-tool", "arguments": {"q": "confidential"}}, + ) + + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=user_api_key_dict) + + assert exc_info.value is guardrail_error + + post_call_failure_hook.assert_awaited_once() + hook_kwargs = post_call_failure_hook.await_args.kwargs + assert hook_kwargs["original_exception"] is guardrail_error + assert hook_kwargs["user_api_key_dict"] is user_api_key_dict + assert hook_kwargs["route"] == "/mcp/call_tool" + request_data = hook_kwargs["request_data"] + assert "raw_headers" not in request_data + assert "mcp_server_auth_headers" not in request_data + standard_logging_object = request_data["litellm_logging_obj"].model_call_details["standard_logging_object"] + assert standard_logging_object["status"] == "failure" + assert standard_logging_object["error_str"] == str(guardrail_error) + + async def test_failure_logging_error_does_not_replace_guardrail_error(self, monkeypatch): + from litellm.proxy import proxy_server + + guardrail_error = HTTPException(status_code=400, detail={"error": "Content blocked"}) + + async def fake_add_litellm_data_to_request(**kwargs): + return kwargs.get("data", {}) + + async def blocking_pre_call_hook(user_api_key_dict, data, call_type): + raise guardrail_error + + failure_logging = AsyncMock(side_effect=RuntimeError("spend log db down")) + monkeypatch.setattr( + proxy_server, "add_litellm_data_to_request", fake_add_litellm_data_to_request, raising=False + ) + monkeypatch.setattr(proxy_server, "proxy_config", {}, raising=False) + monkeypatch.setattr(proxy_server.proxy_logging_obj, "pre_call_hook", blocking_pre_call_hook) + monkeypatch.setattr(rest_endpoints, "fire_mcp_tool_call_failure_logging", failure_logging, raising=False) + + request = _build_request( + path="/mcp-rest/tools/call", + method="POST", + json_body={"server_id": "server-1", "name": "demo-tool", "arguments": {"q": "confidential"}}, + ) + + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.call_tool_rest_api( + request, user_api_key_dict=UserAPIKeyAuth(api_key="hashed-key", request_route="/mcp-rest/tools/call") + ) + + assert exc_info.value is guardrail_error + failure_logging.assert_awaited_once() + async def test_success_logging_cancellation_propagates(self, monkeypatch): fire_logging = AsyncMock(side_effect=asyncio.CancelledError()) monkeypatch.setattr( @@ -2801,7 +3007,7 @@ class TestCallToolRestAPI: class _FakePreCall: def __init__(self, data): - pass + self.data = data async def common_processing_pre_call_logic(self, **kwargs): return None, MagicMock() @@ -2835,6 +3041,58 @@ class TestCallToolRestAPI: assert exc_info.value.headers is not None assert exc_info.value.headers.get("www-authenticate") == challenge + async def test_virtual_mcp_tool_call_guardrail_block_runs_failure_logging(self, monkeypatch): + """A pre_mcp_call guardrail block on the virtual mcp_tool_call branch must write a failure + spend log, same as the direct tool call branch, and still raise the original error.""" + from litellm.proxy import proxy_server + from litellm.proxy._types import LiteLLM_ObjectPermissionTable + + guardrail_error = HTTPException(status_code=400, detail={"error": "Content blocked"}) + + async def fake_contexts(user_api_key_auth): + return [user_api_key_auth] + + async def fake_add_litellm_data_to_request(**kwargs): + return kwargs.get("data", {}) + + async def blocking_pre_call_hook(user_api_key_dict, data, call_type): + raise guardrail_error + + failure_logging = AsyncMock() + monkeypatch.setattr(rest_endpoints, "build_effective_auth_contexts", fake_contexts, raising=False) + monkeypatch.setattr( + proxy_server, "add_litellm_data_to_request", fake_add_litellm_data_to_request, raising=False + ) + monkeypatch.setattr(proxy_server, "proxy_config", {}, raising=False) + monkeypatch.setattr(proxy_server, "general_settings", {}, raising=False) + monkeypatch.setattr(proxy_server.proxy_logging_obj, "pre_call_hook", blocking_pre_call_hook) + monkeypatch.setattr(rest_endpoints, "fire_mcp_tool_call_failure_logging", failure_logging, raising=False) + + user_api_key_dict = UserAPIKeyAuth( + api_key="hashed-key", + request_route="/mcp-rest/tools/call", + object_permission=LiteLLM_ObjectPermissionTable( + object_permission_id="search-scope", + mcp_tool_search_enabled=True, + ), + ) + request = _build_request( + path="/mcp-rest/tools/call", + method="POST", + json_body={"name": "mcp_tool_call", "arguments": {"tool_name": "x", "arguments": {"q": "confidential"}}}, + ) + + with pytest.raises(HTTPException) as exc_info: + await rest_endpoints.call_tool_rest_api(request, user_api_key_dict=user_api_key_dict) + + assert exc_info.value is guardrail_error + failure_logging.assert_awaited_once() + logging_obj, exception, _start_time, user_api_key_auth, request_data = failure_logging.await_args.args + assert exception is guardrail_error + assert user_api_key_auth is user_api_key_dict + assert logging_obj is request_data.get("litellm_logging_obj") + assert logging_obj is not None + class TestGetToolsForSingleServer: """Test _get_tools_for_single_server with object_permission filtering""" diff --git a/tests/test_litellm/proxy/auth/test_route_checks.py b/tests/test_litellm/proxy/auth/test_route_checks.py index 38a0e85c1ea..c8b3d789665 100644 --- a/tests/test_litellm/proxy/auth/test_route_checks.py +++ b/tests/test_litellm/proxy/auth/test_route_checks.py @@ -3916,3 +3916,28 @@ def test_claude_code_marketplace_routes_open_to_internal_users(route): """Per-skill visibility is enforced inside the handler, so the route gate must let non-admins through.""" assert RouteChecks.is_llm_api_route(route) is True assert _gate(route, LitellmUserRoles.INTERNAL_USER.value) == "allowed" + + +@pytest.mark.parametrize("user_role", [None, LitellmUserRoles.INTERNAL_USER.value, LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value]) +def test_auto_router_session_is_reachable_by_any_key_but_benchmarks_stays_admin_only(user_role): + valid_token = UserAPIKeyAuth(api_key="hash-of-caller", user_role=user_role) + request = MagicMock(spec=Request) + request.query_params = {"session_id": "sess-1"} + + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=None, + _user_role=user_role, + route="/auto_router/session", + request=request, + valid_token=valid_token, + request_data={}, + ) + with pytest.raises(Exception, match="Only proxy admin"): + RouteChecks.non_proxy_admin_allowed_routes_check( + user_obj=None, + _user_role=user_role, + route="/auto_router/benchmarks", + request=request, + valid_token=valid_token, + request_data={}, + ) diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_commands.py b/tests/test_litellm/proxy/client/cli/autoroute/test_commands.py index 77742ea9f9f..028ab58843f 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_commands.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_commands.py @@ -6,6 +6,7 @@ from typing import Optional import yaml from click.testing import CliRunner +from litellm.proxy.client.cli.commands.claude_settings import ClaudeSettingsError from litellm.proxy.client.cli.commands.autoroute import commands as commands_module from litellm.proxy.client.cli.commands.autoroute import process as process_module from litellm.proxy.client.cli.commands.autoroute.commands import down, up @@ -163,6 +164,7 @@ class TestUpCommand: # `lite configure claude --model` or a user pin would 400 on the first message. assert captured["settings"]["model"] == "autorouter" assert captured["settings"]["env"]["ANTHROPIC_DEFAULT_SONNET_MODEL"] == "autorouter" + assert captured["settings"]["statusLine"]["command"].endswith("statusline.py") assert captured["settings_mode"] == 0o600 assert terminate_calls == [99999] @@ -253,6 +255,33 @@ class TestUpCommand: assert not pid_record_path.exists() assert not backup_path.exists() + def test_a_status_line_install_failure_leaves_no_backup_behind(self, monkeypatch, tmp_path): + # The install runs before the backup is written, so a failure cannot strand a backup that + # would make every later `lite configure` / `lite autoroute up` think a session still owns settings.json + config_path, _log_path, claude_settings_path, backup_path, pid_record_path = _patch_paths(monkeypatch, tmp_path) + config_path.write_text(yaml.safe_dump({"model_list": []})) + claude_settings_path.write_text(json.dumps({"theme": "dark"})) + + def boom(): + raise ClaudeSettingsError("disk full") + + fake_process = FakeProcess(pid=778) + terminate_calls = [] + monkeypatch.setattr(commands_module, "launch_proxy", lambda *a, **k: fake_process) + monkeypatch.setattr(commands_module, "poll_liveliness", lambda *a, **k: None) + monkeypatch.setattr(commands_module, "is_port_available", lambda port: True) + monkeypatch.setattr(commands_module, "terminate", lambda pid, **k: terminate_calls.append(pid)) + monkeypatch.setattr(commands_module, "install_statusline_script", boom) + monkeypatch.setattr(commands_module.secrets, "token_urlsafe", lambda n: "fixed-master-key") + + result = self.runner.invoke(up) + + assert result.exit_code != 0 and "disk full" in result.output + assert terminate_calls == [778] + assert not pid_record_path.exists() + assert not backup_path.exists() + assert json.loads(claude_settings_path.read_text()) == {"theme": "dark"} + def test_up_uses_the_same_port_and_master_key_across_runs(self, monkeypatch, tmp_path): """The LIT-4607/LIT-4608 regression: a client configured against one session must keep working in the next, so consecutive runs must patch settings with an identical base URL diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_config.py b/tests/test_litellm/proxy/client/cli/autoroute/test_config.py index f399b6f957a..47b5459d489 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_config.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_config.py @@ -161,7 +161,13 @@ class TestBuildGeneratedModelList: config = _base_config(classifier=HeuristicClassifier(), semantic_matching=NoSemanticMatching()) autorouter = next(m for m in build_generated_model_list(config) if m["model_name"] == "autorouter") router_config = autorouter["litellm_params"]["complexity_router_config"] - assert set(router_config.keys()) == {"tiers", "default_model"} + assert set(router_config.keys()) == {"tiers", "default_model", "return_raw_model_name"} + + def test_the_generated_router_reports_the_tier_model_it_routed_to(self): + # The status line reads the routed model from the response body, which the proxy restamps to the + # requested alias unless the deployment opts out; "autorouter" on every line would tell nothing. + autorouter = next(m for m in build_generated_model_list(_base_config()) if m["model_name"] == "autorouter") + assert autorouter["litellm_params"]["complexity_router_config"]["return_raw_model_name"] is True class TestBuildGeneratedProxyConfig: diff --git a/tests/test_litellm/proxy/client/cli/conftest.py b/tests/test_litellm/proxy/client/cli/conftest.py index 50c76d3f125..c77f516a768 100644 --- a/tests/test_litellm/proxy/client/cli/conftest.py +++ b/tests/test_litellm/proxy/client/cli/conftest.py @@ -5,6 +5,8 @@ from typing import Final import pytest +from litellm.proxy.client.cli.commands import claude_settings + REAL_CLAUDE_SETTINGS: Final = Path(os.path.expanduser("~")) / ".claude" / "settings.json" @@ -12,6 +14,11 @@ def _current_bytes() -> bytes | None: return REAL_CLAUDE_SETTINGS.read_bytes() if REAL_CLAUDE_SETTINGS.exists() else None +@pytest.fixture(autouse=True) +def _statusline_script_under_tmp(monkeypatch: pytest.MonkeyPatch, tmp_path: Path) -> None: + monkeypatch.setattr(claude_settings, "STATUSLINE_SCRIPT_PATH", tmp_path / "litellm-home" / "statusline.py") + + @pytest.fixture(autouse=True) def isolated_claude_home(tmp_path: Path, monkeypatch: pytest.MonkeyPatch) -> Iterator[Path]: before: Final = _current_bytes() diff --git a/tests/test_litellm/proxy/client/cli/test_agents.py b/tests/test_litellm/proxy/client/cli/test_agents.py index bebea285edc..7804435a60d 100644 --- a/tests/test_litellm/proxy/client/cli/test_agents.py +++ b/tests/test_litellm/proxy/client/cli/test_agents.py @@ -164,27 +164,6 @@ class TestBuildAgentEnv: assert env["PATH"] == "/usr/bin" assert base == {"PATH": "/usr/bin", "ANTHROPIC_API_KEY": "real-key"} - def test_anthropic_profile_leaves_the_bearer_to_the_api_key_helper(self): - env = build_agent_env( - {"ANTHROPIC_AUTH_TOKEN": "stale-token", "ANTHROPIC_API_KEY": "real-key"}, - "http://localhost:4000/", - "sk-key", - frozenset({"anthropic"}), - export_anthropic_token=False, - ) - assert "ANTHROPIC_AUTH_TOKEN" not in env - assert "ANTHROPIC_API_KEY" not in env - assert env["ANTHROPIC_BASE_URL"] == "http://localhost:4000" - assert env["ENABLE_TOOL_SEARCH"] == "true" - assert env["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1" - - def test_helper_mode_still_exports_the_openai_key(self): - env = build_agent_env( - {}, "http://localhost:4000", "sk-key", frozenset({"anthropic", "openai"}), export_anthropic_token=False - ) - assert "ANTHROPIC_AUTH_TOKEN" not in env - assert env["OPENAI_API_KEY"] == "sk-key" - class TestAgentLaunchArgs: def test_claude_and_opencode_get_no_extra_args(self): @@ -531,25 +510,6 @@ class TestRunAgent: assert "ANTHROPIC_API_KEY" not in env assert "OPENAI_BASE_URL" not in env - def test_helper_supplied_token_never_reaches_the_launch_env(self): - calls = {} - verified = [] - - run_agent( - "http://localhost:4000", - "sk-key", - ["claude"], - base_env={"PATH": "/usr/bin", "ANTHROPIC_AUTH_TOKEN": "stale-token"}, - which=lambda name: "/usr/local/bin/claude", - verify=lambda base_url, api_key: verified.append(api_key), - launcher=lambda p, a, e: calls.update(env=dict(e)), - export_anthropic_token=False, - ) - - assert verified == ["sk-key"] - assert "ANTHROPIC_AUTH_TOKEN" not in calls["env"] - assert calls["env"]["ANTHROPIC_BASE_URL"] == "http://localhost:4000" - def test_codex_gets_openai_env(self): calls = {} run_agent( @@ -1093,101 +1053,6 @@ class TestAgentCommands: in result.output ) - def _invoke_claude_with_settings(self, tmp_path, settings, obj, *, default_settings=None): - config_dir = tmp_path / "claude-config" - config_dir.mkdir() - if settings is not None: - (config_dir / "settings.json").write_text(json.dumps(settings)) - default_path = tmp_path / "home-claude" / "settings.json" - default_path.parent.mkdir() - if default_settings is not None: - default_path.write_text(json.dumps(default_settings)) - captured = {} - with ( - patch(f"{CLAUDE_SETTINGS_MODULE}.CLAUDE_SETTINGS_PATH", default_path), - patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value="/usr/local/bin/lite"), - patch(f"{AGENTS_MODULE}.run_agent", side_effect=lambda b, k, c, **kw: captured.update(kw)), - ): - result = self.runner.invoke( - _agent_command("claude"), [], obj=obj, env={"CLAUDE_CONFIG_DIR": str(config_dir)} - ) - assert result.exit_code == 0, result.output - return captured, result.output - - def test_helper_is_read_from_the_config_dir_claude_code_uses(self, tmp_path): - captured, output = self._invoke_claude_with_settings( - tmp_path, - {"apiKeyHelper": "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"}, - {"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True}, - ) - - assert captured["export_anthropic_token"] is False - assert str(tmp_path / "claude-config" / "settings.json") in output - - def test_helper_only_in_the_default_file_keeps_the_env_token_when_config_dir_points_elsewhere(self, tmp_path): - captured, output = self._invoke_claude_with_settings( - tmp_path, - None, - {"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True}, - default_settings={"apiKeyHelper": "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"}, - ) - - assert captured["export_anthropic_token"] is True - assert "apiKeyHelper" not in output - - def test_stored_login_with_a_matching_helper_leaves_the_token_to_the_helper(self, tmp_path): - captured, output = self._invoke_claude_with_settings( - tmp_path, - {"apiKeyHelper": "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"}, - {"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True}, - ) - - assert captured["export_anthropic_token"] is False - assert "reads its key from the apiKeyHelper" in output - - def test_explicit_key_is_exported_even_when_a_helper_matches(self, tmp_path): - captured, output = self._invoke_claude_with_settings( - tmp_path, - {"apiKeyHelper": "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token"}, - {"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": False}, - ) - - assert captured["export_anthropic_token"] is True - assert "apiKeyHelper" not in output - - def test_helper_for_another_proxy_keeps_the_env_token(self, tmp_path): - captured, _ = self._invoke_claude_with_settings( - tmp_path, - {"apiKeyHelper": "/usr/local/bin/lite --base-url https://other.example.com auth print-token"}, - {"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True}, - ) - - assert captured["export_anthropic_token"] is True - - def test_no_claude_settings_keeps_the_env_token(self, tmp_path): - captured, _ = self._invoke_claude_with_settings( - tmp_path, - None, - {"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True}, - ) - - assert captured["export_anthropic_token"] is True - - def test_codex_never_consults_claude_settings(self): - captured = {} - with ( - patch(f"{AGENTS_MODULE}.lite_api_key_helper_configured", side_effect=AssertionError("consulted")), - patch(f"{AGENTS_MODULE}.run_agent", side_effect=lambda b, k, c, **kw: captured.update(kw)), - ): - result = self.runner.invoke( - _agent_command("codex"), - [], - obj={"base_url": "http://localhost:4000", "api_key": "sk-key", "api_key_from_token_file": True}, - ) - - assert result.exit_code == 0, result.output - assert captured["export_anthropic_token"] is True - def test_codex_shows_friendly_name(self): captured = {} with patch( @@ -1338,3 +1203,53 @@ class TestAgentCommands: ) assert result.exit_code == 0, result.output assert captured["reattach_terminal"] is None + + +class TestPrepareCodex: + def test_registers_the_installed_script_as_a_session_scoped_stop_hook(self): + from litellm.proxy.client.cli.commands.agents import prepare_codex + + args = prepare_codex("http://localhost:4000", "sk-key", {}, install=lambda: "/py /home/me/.litellm/statusline.py") + assert args == ( + "-c", + 'hooks.Stop=[{hooks=[{type="command",command="/py /home/me/.litellm/statusline.py"}]}]', + ) + + def test_a_failed_install_is_an_agent_error_not_a_crash(self): + from litellm.proxy.client.cli.commands.agents import AgentRunError, prepare_codex + from litellm.proxy.client.cli.commands.claude_settings import ClaudeSettingsError + + def boom(): + raise ClaudeSettingsError("disk full") + + with pytest.raises(AgentRunError, match="disk full"): + prepare_codex("http://localhost:4000", "sk-key", {}, install=boom) + + def test_a_config_that_already_declares_hooks_keeps_them_and_skips_ours(self, tmp_path): + from litellm.proxy.client.cli.commands.agents import prepare_codex + + warnings = [] + env = {"CODEX_HOME": str(tmp_path)} + for body in ('[[hooks.Stop]]\nhooks = [{ type = "command", command = "mine" }]\n', 'hooks.Stop = []\n', "[hooks]\n"): + (tmp_path / "config.toml").write_text(body) + assert prepare_codex("http://localhost:4000", "sk", env, install=lambda: "/py /s.py", warn=warnings.append) == () + (tmp_path / "config.toml").write_text('model = "gpt-5.6-sol"\n[projects."/x"]\ntrust_level = "trusted"\n') + assert prepare_codex("http://localhost:4000", "sk", env, install=lambda: "/py /s.py", warn=warnings.append) != () + assert len(warnings) == 3 and "already declares hooks" in warnings[0] + + def test_a_config_that_cannot_be_read_or_decoded_still_lets_codex_launch(self, tmp_path): + # A UTF-16 config.toml (a Windows Notepad save) is Codex's problem to report at launch, not a reason + # for the hook pre-check to abort `lite codex` with a traceback before Codex ever starts. + from litellm.proxy.client.cli.commands.agents import codex_declares_stop_hooks, prepare_codex + + config = tmp_path / "config.toml" + config.write_bytes('[[hooks.Stop]]\nhooks = [{ type = "command", command = "mine" }]\n'.encode("utf-16")) + assert codex_declares_stop_hooks(config) is False + assert codex_declares_stop_hooks(tmp_path / "absent.toml") is False + args = prepare_codex("http://localhost:4000", "sk", {"CODEX_HOME": str(tmp_path)}, install=lambda: "/py /s.py") + assert args[0] == "-c" and "hooks.Stop=" in args[1] + + def test_codex_is_wired_through_the_preparer_registry(self): + from litellm.proxy.client.cli.commands.agents import _PREPARERS, prepare_codex + + assert _PREPARERS["codex"] is prepare_codex diff --git a/tests/test_litellm/proxy/client/cli/test_auth_commands.py b/tests/test_litellm/proxy/client/cli/test_auth_commands.py index 1f314e0c9d8..3a7792db1fe 100644 --- a/tests/test_litellm/proxy/client/cli/test_auth_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_auth_commands.py @@ -17,8 +17,15 @@ from litellm.litellm_core_utils.cli_keyring import ( SecretErased, SecretStored, ) -from litellm.litellm_core_utils.cli_token_utils import CliTokenRecord, save_cli_token +from litellm.litellm_core_utils.cli_token_utils import ( + CliTokenRecord, + CredentialNotRecorded, + CredentialNotSaved, + save_cli_token, +) from litellm.proxy.client.cli import cli +from litellm.proxy.client.cli.commands import claude_settings as claude_settings_module +from litellm.proxy.client.cli.commands.claude_settings import SettingsFileOwner from litellm.proxy.client.cli.commands.auth import ( get_stored_api_key, login, @@ -26,8 +33,6 @@ from litellm.proxy.client.cli.commands.auth import ( print_token, whoami, ) -from litellm.proxy.client.cli.commands import claude_settings as claude_settings_module -from litellm.proxy.client.cli.commands.claude_settings import SettingsFileOwner @pytest.fixture @@ -1394,7 +1399,9 @@ class TestLoginConfigClaude: monkeypatch.setattr(claude_settings_module, "CONFIGURE_STATE_PATH", tmp_path / "claude_configure_state.json") return backup_path - def _run_login(self, tmp_path, monkeypatch, args, base_url="https://test.example.com", *, config_dir_env=None): + def _run_login( + self, tmp_path, monkeypatch, args, base_url="https://test.example.com", *, config_dir_env=None, stored=None + ): settings_path = tmp_path / "claude" / "settings.json" backup_path = self._isolate_default_settings(tmp_path, monkeypatch) env = {"CLAUDE_CONFIG_DIR": str(settings_path.parent)} if config_dir_env is None else config_dir_env @@ -1411,12 +1418,8 @@ class TestLoginConfigClaude: patch("webbrowser.open"), patch("requests.post", return_value=_mock_cli_sso_start_response()), patch("requests.get", return_value=poll_response), - patch("litellm.proxy.client.cli.commands.auth.save_cli_token"), + patch("litellm.proxy.client.cli.commands.auth.save_cli_token", return_value=stored or SecretStored()), patch("litellm.proxy.client.cli.interface.show_commands"), - patch( - "litellm.proxy.client.cli.commands.claude_settings.shutil.which", - return_value="/usr/local/bin/lite", - ), ): result = self.runner.invoke(login, args, obj={"base_url": base_url}, env=env) return result, settings_path, backup_path @@ -1436,10 +1439,13 @@ class TestLoginConfigClaude: written = json.loads(settings_path.read_text()) assert written["env"]["ANTHROPIC_BASE_URL"] == "https://test.example.com" assert written["env"]["ENABLE_TOOL_SEARCH"] == "true" - assert written["apiKeyHelper"] == "/usr/local/bin/lite --base-url https://test.example.com auth print-token" + # The minted key goes in as a static token: an apiKeyHelper would make Claude Code spawn `lite` (and + # its keychain probe) on every credential refresh, which is what this flag used to write. + assert written["env"]["ANTHROPIC_AUTH_TOKEN"] == "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt" + assert "apiKeyHelper" not in written assert f"Configured Claude Code: {settings_path} now routes through https://test.example.com." in result.output - assert "pins a proxy model for every tier" not in result.output - assert "the model Claude Code starts on" in result.output + assert "run `lite login --config-claude` again after it expires" in result.output + assert "the model Claude Code starts and resumes on" in result.output def test_flag_preserves_unrelated_settings_on_an_existing_file(self, tmp_path, monkeypatch): settings_path = tmp_path / "claude" / "settings.json" @@ -1490,7 +1496,7 @@ class TestLoginConfigClaude: assert result.exit_code == 0, result.output written = json.loads(settings_path.read_text()) - assert written["apiKeyHelper"] == "/usr/local/bin/lite --base-url https://test.example.com auth print-token" + assert written["env"]["ANTHROPIC_AUTH_TOKEN"] == "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt" assert f"Configured Claude Code: {settings_path} now routes through https://test.example.com." in result.output def test_flag_keeps_a_config_dir_receipt_apart_from_the_default_file_receipt(self, tmp_path, monkeypatch): @@ -1503,6 +1509,42 @@ class TestLoginConfigClaude: assert len(receipts) == 1 assert json.loads(receipts[0].read_text())["file_existed"] is False + def test_a_second_login_replaces_the_key_and_unconfigure_still_restores_the_original(self, tmp_path, monkeypatch): + # The stored key expires daily, so the flag is re-run per login; the receipt must keep owning the + # slot across re-logins and hand back what was there before the first one. + from litellm.proxy.client.cli.commands.configure import unconfigure_claude + + settings_path = tmp_path / "claude" / "settings.json" + settings_path.parent.mkdir(parents=True) + settings_path.write_text(json.dumps({"theme": "dark", "env": {"ANTHROPIC_AUTH_TOKEN": "sk-theirs"}})) + self._run_login(tmp_path, monkeypatch, ["--config-claude"]) + first = json.loads(settings_path.read_text())["env"]["ANTHROPIC_AUTH_TOKEN"] + self._run_login(tmp_path, monkeypatch, ["--config-claude"]) + assert json.loads(settings_path.read_text())["env"]["ANTHROPIC_AUTH_TOKEN"] == first != "sk-theirs" + + result = self.runner.invoke(unconfigure_claude, [], env={"CLAUDE_CONFIG_DIR": str(settings_path.parent)}) + assert result.exit_code == 0, result.output + assert json.loads(settings_path.read_text()) == {"theme": "dark", "env": {"ANTHROPIC_AUTH_TOKEN": "sk-theirs"}} + + @pytest.mark.parametrize( + "stored", + [CredentialNotSaved("read-only ~/.litellm"), CredentialNotRecorded()], + ids=["nothing-kept-it", "keychain-took-it-file-refused"], + ) + def test_claude_code_is_configured_even_when_the_cli_could_not_keep_the_credential( + self, tmp_path, monkeypatch, stored + ): + # The key is in hand either way, and --config-claude asked for exactly that key to be written into + # settings.json; whether the CLI's own token file or keychain kept a copy is a separate outcome. + result, settings_path, _backup_path = self._run_login(tmp_path, monkeypatch, ["--config-claude"], stored=stored) + + assert result.exit_code == 0, result.output + written = json.loads(settings_path.read_text()) + assert written["env"]["ANTHROPIC_AUTH_TOKEN"] == "eyJhbGciOiJIUzI1NiIsInR5cCI6IkpXVCJ9.test.jwt" + assert f"Configured Claude Code: {settings_path}" in result.output + assert "even though the CLI itself could not keep it" in result.output + assert "You can now use the CLI without specifying --api-key" not in result.output + def test_settings_failure_is_reported_without_claiming_login_failed(self, tmp_path, monkeypatch): settings_path = tmp_path / "claude" / "settings.json" settings_path.parent.mkdir(parents=True) diff --git a/tests/test_litellm/proxy/client/cli/test_claude_settings.py b/tests/test_litellm/proxy/client/cli/test_claude_settings.py index fc2d98f2264..cf52d41e963 100644 --- a/tests/test_litellm/proxy/client/cli/test_claude_settings.py +++ b/tests/test_litellm/proxy/client/cli/test_claude_settings.py @@ -1,7 +1,9 @@ import json import os +import pathlib import shlex import stat +import sys import time from pathlib import Path from unittest.mock import patch @@ -9,8 +11,6 @@ from unittest.mock import patch import pytest from click.testing import CliRunner -from litellm.litellm_core_utils.cli_token_utils import CliTokenRecord -from litellm.proxy.client.cli import cli from litellm.litellm_core_utils.private_json import commit_staged_json from litellm.proxy.client.cli.commands.claude_settings import ( ANTHROPIC_DEFAULT_MODEL_ENV_KEYS, @@ -21,7 +21,6 @@ from litellm.proxy.client.cli.commands.claude_settings import ( OWNED_ENV_KEYS, OWNED_TOP_LEVEL_KEYS, SETTINGS_FILE_OWNERS, - ApiKeyHelper, ClaudeSettingsError, KeepModel, SettingsFileOwner, @@ -30,11 +29,12 @@ from litellm.proxy.client.cli.commands.claude_settings import ( UnpinModel, claude_settings_path, configure_claude_settings, + install_statusline_script, configure_state_path, - lite_api_key_helper_configured, merge_claude_settings, - resolve_api_key_helper, + statusline_command, unconfigure_claude_settings, + with_status_line, ) @@ -45,101 +45,33 @@ def _owners(*backup_paths): CLAUDE_SETTINGS_MODULE = "litellm.proxy.client.cli.commands.claude_settings" AUTH_MODULE = "litellm.proxy.client.cli.commands.auth" -WINDOWS_LITE_EXE = "C:\\Users\\u\\AppData\\Local\\Programs\\Python\\Python313\\Scripts\\lite.EXE" - -CMD_METACHARACTERS = frozenset("&|<>^()") -CMD_PERCENT_GUARD = "%%cd:~,%" - - -def _through_cmd_exe(command): - """The line cmd.exe hands to CreateProcess after reading the apiKeyHelper. - - A `"` toggles cmd's quote state and the metacharacters only act outside it. cmd expands - `%VAR%` even inside quotes, so every `%` has to arrive as the `%%cd:~,%` guard: the first - `%` has no variable name and stays literal, and `%cd:~,%` is a zero length substring of `cd`. - """ - assert not any(CMD_METACHARACTERS & set(run) for run in command.split('"')[::2]), command - assert command.count("%") == 3 * command.count(CMD_PERCENT_GUARD), command - return command.replace(CMD_PERCENT_GUARD, "%") - - -def _through_c_runtime(command_line): - """argv as the Microsoft C runtime builds it for the `lite` executable. - - Outside quotes whitespace ends an argument. A `"` toggles quoting, and inside quotes `""` - is a literal quote. Backslashes are literal unless they run up to a `"`, where each pair - is one backslash and an odd one left over makes the quote literal. - """ - argv = [] - current = None - quoted = False - i = 0 - while i < len(command_line): - ch = command_line[i] - if ch in " \t" and not quoted: - if current is not None: - argv.append(current) - current = None - i += 1 - continue - if current is None: - current = "" - if ch == "\\": - run = len(command_line[i:]) - len(command_line[i:].lstrip("\\")) - before_quote = command_line[i + run : i + run + 1] == '"' - current += "\\" * (run // 2 if before_quote else run) - if before_quote and run % 2: - current += '"' - i += 1 - i += run - elif ch == '"': - if quoted and command_line[i + 1 : i + 2] == '"': - current += '"' - i += 1 - else: - quoted = not quoted - i += 1 - else: - current += ch - i += 1 - return argv if current is None else [*argv, current] - - @pytest.fixture def paths(tmp_path): return tmp_path / "claude" / "settings.json", tmp_path / "backup.json" -@pytest.fixture -def lite_on_path(): - with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value="/usr/local/bin/lite"): - yield - - -def _helper_configure(base_url, settings_path, owners, state_path=None): - """`lite login --config-claude`'s shape: the login credential behind apiKeyHelper, no pinned model.""" +def _static_configure(base_url, settings_path, owners, state_path=None): + """`lite configure claude --api-key`'s shape: a virtual key as a static token, no pinned model.""" state = state_path if state_path is not None else settings_path.parent.parent / "state.json" - root = base_url.rstrip("/") - configure_claude_settings( - root, ApiKeyHelper(resolve_api_key_helper(root)), KeepModel(), settings_path, state, owners - ) + configure_claude_settings(base_url.rstrip("/"), StaticToken("sk-virtual-key"), KeepModel(), settings_path, state, owners) -class TestConfigureWithTheLoginHelper: - def test_creates_the_file_and_its_parent_when_missing(self, paths, lite_on_path): +class TestConfigureClaudeSettings: + def test_creates_the_file_and_its_parent_when_missing(self, paths): settings_path, backup_path = paths assert not settings_path.parent.exists() - _helper_configure("https://proxy.example.com/", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com/", settings_path, _owners(backup_path)) written = json.loads(settings_path.read_text()) assert written["env"]["ANTHROPIC_BASE_URL"] == "https://proxy.example.com" + assert written["env"]["ANTHROPIC_AUTH_TOKEN"] == "sk-virtual-key" assert written["env"]["ENABLE_TOOL_SEARCH"] == "true" assert written["env"]["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1" - assert written["apiKeyHelper"] == "/usr/local/bin/lite --base-url https://proxy.example.com auth print-token" - assert "model" not in written + assert "apiKeyHelper" not in written + assert "model" not in written and "ANTHROPIC_MODEL" not in written["env"] - def test_updates_an_existing_file_preserving_unrelated_settings(self, paths, lite_on_path): + def test_updates_an_existing_file_preserving_unrelated_settings(self, paths): settings_path, backup_path = paths settings_path.parent.mkdir(parents=True) settings_path.write_text( @@ -148,183 +80,87 @@ class TestConfigureWithTheLoginHelper: "theme": "dark", "permissions": {"allow": ["Bash"]}, "env": {"SOME_OTHER_VAR": "keep-me", "ANTHROPIC_BASE_URL": "https://old.example.com"}, - "apiKeyHelper": "old-helper", } ) ) - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) written = json.loads(settings_path.read_text()) assert written["theme"] == "dark" assert written["permissions"] == {"allow": ["Bash"]} assert written["env"]["SOME_OTHER_VAR"] == "keep-me" assert written["env"]["ANTHROPIC_BASE_URL"] == "https://proxy.example.com" - assert written["apiKeyHelper"] != "old-helper" - def test_rerunning_against_a_new_proxy_refreshes_both_base_url_and_helper(self, paths, lite_on_path): - settings_path, backup_path = paths - - _helper_configure("https://first.example.com", settings_path, _owners(backup_path)) - _helper_configure("https://second.example.com", settings_path, _owners(backup_path)) - - written = json.loads(settings_path.read_text()) - assert written["env"]["ANTHROPIC_BASE_URL"] == "https://second.example.com" - assert "second.example.com" in written["apiKeyHelper"] - assert "first.example.com" not in written["apiKeyHelper"] - - def test_drops_stray_static_credentials_so_the_helper_token_wins(self, paths, lite_on_path): - # Claude Code prefers ANTHROPIC_AUTH_TOKEN over apiKeyHelper, so a virtual key left behind - # by an earlier `lite configure claude --api-key` would silently keep winning. + def test_a_helper_left_by_an_older_lite_is_stripped_so_only_the_static_token_is_sent(self, paths): + # Older `lite` versions wrote `apiKeyHelper: lite auth print-token`; Claude Code would keep spawning + # `lite` (and its keychain probe) on every credential refresh, so configure takes the slot over. settings_path, backup_path = paths settings_path.parent.mkdir(parents=True) settings_path.write_text( - json.dumps({"env": {"ANTHROPIC_API_KEY": "sk-leaked", "ANTHROPIC_AUTH_TOKEN": "sk-old"}}) + json.dumps({"apiKeyHelper": "/usr/local/bin/lite auth print-token", "env": {"ANTHROPIC_API_KEY": "sk-leaked"}}) ) - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) - env = json.loads(settings_path.read_text())["env"] - assert "ANTHROPIC_API_KEY" not in env and "ANTHROPIC_AUTH_TOKEN" not in env + written = json.loads(settings_path.read_text()) + assert "apiKeyHelper" not in written + assert "ANTHROPIC_API_KEY" not in written["env"] + assert written["env"]["ANTHROPIC_AUTH_TOKEN"] == "sk-virtual-key" - def test_written_file_is_owner_only(self, paths, lite_on_path): + def test_written_file_is_owner_only(self, paths): settings_path, backup_path = paths - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) assert stat.S_IMODE(settings_path.stat().st_mode) == 0o600 - def test_refuses_while_lite_up_holds_a_backup(self, paths, lite_on_path): + def test_refuses_while_lite_up_holds_a_backup(self, paths): settings_path, backup_path = paths backup_path.write_text("{}") with pytest.raises(ClaudeSettingsError, match="lite down"): - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) assert not settings_path.exists() - def test_refuses_on_corrupt_existing_settings_without_touching_the_file(self, paths, lite_on_path): + def test_refuses_on_corrupt_existing_settings_without_touching_the_file(self, paths): settings_path, backup_path = paths settings_path.parent.mkdir(parents=True) settings_path.write_text("not json at all {{{") with pytest.raises(ClaudeSettingsError, match="invalid JSON"): - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) assert settings_path.read_text() == "not json at all {{{" - def test_reports_an_actionable_error_when_lite_is_not_on_path(self, paths): - settings_path, backup_path = paths - with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value=None): - with pytest.raises(ClaudeSettingsError, match="Could not find `lite`"): - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) - - assert not settings_path.exists() - - def test_reports_an_actionable_error_on_a_non_utf8_file(self, paths, lite_on_path): - """Bytes that are not valid UTF-8 must not escape as UnicodeDecodeError. - - UnicodeDecodeError is a ValueError, not an OSError, so a decode-side catch - is easy to miss; login's broad `except Exception` would then relabel it as - an authentication failure and exit 0. - """ + def test_reports_an_actionable_error_on_a_non_utf8_file(self, paths): + # UnicodeDecodeError is a ValueError, not an OSError, so a decode-side catch is easy to miss. settings_path, backup_path = paths settings_path.parent.mkdir(parents=True) settings_path.write_bytes(b'{"theme": "\xff\xfe"}') with pytest.raises(ClaudeSettingsError, match="invalid JSON"): - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) - def test_reports_an_actionable_error_when_the_file_cannot_be_read(self, paths, lite_on_path): - """An unreadable settings file must not surface as "Authentication failed". - - login wraps the whole flow in a broad `except Exception`, so any OSError - escaping this function gets relabelled as an auth failure and sends the - user looking at their SSO config instead of at file permissions. - """ + def test_reports_an_actionable_error_when_the_file_cannot_be_read(self, paths): settings_path, backup_path = paths settings_path.parent.mkdir(parents=True) settings_path.mkdir() with pytest.raises(ClaudeSettingsError, match="Could not read"): - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) - def test_reports_an_actionable_error_when_the_file_cannot_be_written(self, paths, lite_on_path): + def test_reports_an_actionable_error_when_the_file_cannot_be_written(self, paths): settings_path, backup_path = paths settings_path.parent.mkdir(parents=True) settings_path.parent.chmod(0o500) try: with pytest.raises(ClaudeSettingsError, match="Could not write"): - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) finally: settings_path.parent.chmod(0o700) assert not settings_path.exists() -class TestApiKeyHelperIsActuallyInvocable: - """The helper string is executed verbatim by Claude Code, so it has to parse. - - Asserting only on its text is what let a malformed command (`--base-url`, a - top-level group option, placed after the `print-token` subcommand) ship: click - rejects it with "No such option" and every Claude Code request loses its token. - """ - - def _helper_args(self, base_url): - with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value="/usr/local/bin/lite"): - return shlex.split(resolve_api_key_helper(base_url))[1:] - - def test_the_generated_command_parses(self): - result = CliRunner().invoke(cli, self._helper_args("http://localhost:4000")) - - assert "No such option" not in result.output - assert result.exit_code != 2 - - def test_the_generated_command_reaches_print_token(self): - with patch(f"{AUTH_MODULE}.load_cli_token", return_value=None): - result = CliRunner().invoke(cli, self._helper_args("http://localhost:4000")) - - assert "Not authenticated" in result.output - - def test_the_generated_command_carries_the_base_url_through(self): - stale = CliTokenRecord( - base_url="http://other-proxy.example.com", - key="sk-stale", - timestamp=time.time(), - ) - with patch(f"{AUTH_MODULE}.load_cli_token", return_value=stale): - result = CliRunner().invoke(cli, self._helper_args("http://localhost:4000")) - - assert "Not authenticated for this server" in result.output - - def _windows_argv(self, lite_exe, base_url): - with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value=lite_exe): - helper = resolve_api_key_helper(base_url, platform="win32") - return _through_c_runtime(_through_cmd_exe(helper)) - - @pytest.mark.parametrize( - ("lite_exe", "base_url"), - [ - (WINDOWS_LITE_EXE, "http://localhost:4000"), - ("C:\\Program Files\\LiteLLM\\lite.EXE", "https://gateway.example.com/?a=1&b=2"), - ("C:\\Users\\u\\Scripts\\lite.EXE", "https://gateway.example.com/team%20a/%7Eproxy"), - ('C:\\odd "dir"\\lite.EXE', "http://localhost:4000/x\\"), - ], - ) - def test_the_windows_command_survives_cmd_exe_and_the_c_runtime(self, lite_exe, base_url): - assert self._windows_argv(lite_exe, base_url) == [lite_exe, "--base-url", base_url, "auth", "print-token"] - - def test_the_windows_command_carries_the_base_url_through_cmd_quoting(self): - stale = CliTokenRecord( - base_url="http://other-proxy.example.com", - key="sk-stale", - timestamp=time.time(), - ) - argv = self._windows_argv(WINDOWS_LITE_EXE, "http://localhost:4000") - with patch(f"{AUTH_MODULE}.load_cli_token", return_value=stale): - result = CliRunner().invoke(cli, argv[1:]) - - assert argv[0] == WINDOWS_LITE_EXE - assert "Not authenticated for this server" in result.output - - class TestConflictingOwnersOfTheSettingsFile: """Both `lite up` and `lite autoroute up` restore a backup when they stop. @@ -332,7 +168,7 @@ class TestConflictingOwnersOfTheSettingsFile: write, which is the exact hazard the guard exists to prevent. """ - def test_any_owner_holding_a_backup_blocks_the_write(self, tmp_path, lite_on_path): + def test_any_owner_holding_a_backup_blocks_the_write(self, tmp_path): settings_path = tmp_path / "claude" / "settings.json" for index, owner in enumerate(SETTINGS_FILE_OWNERS): @@ -340,20 +176,20 @@ class TestConflictingOwnersOfTheSettingsFile: backup.write_text("{}") stand_in = SettingsFileOwner(backup, owner.start_command, owner.stop_command) with pytest.raises(ClaudeSettingsError, match="currently managing"): - _helper_configure("https://proxy.example.com", settings_path, (stand_in,)) + _static_configure("https://proxy.example.com", settings_path, (stand_in,)) backup.unlink() assert not settings_path.exists() - def test_the_error_names_the_owner_that_actually_holds_the_file(self, tmp_path, lite_on_path): + def test_the_error_names_the_owner_that_actually_holds_the_file(self, tmp_path): settings_path = tmp_path / "claude" / "settings.json" backup = tmp_path / "auto.json" backup.write_text("{}") autoroute = SettingsFileOwner(backup, "lite autoroute up", "lite autoroute down") with pytest.raises(ClaudeSettingsError, match="`lite autoroute up` is currently managing"): - _helper_configure("https://proxy.example.com", settings_path, (autoroute,)) + _static_configure("https://proxy.example.com", settings_path, (autoroute,)) with pytest.raises(ClaudeSettingsError, match="Run `lite autoroute down` first"): - _helper_configure("https://proxy.example.com", settings_path, (autoroute,)) + _static_configure("https://proxy.example.com", settings_path, (autoroute,)) def test_the_registry_matches_the_paths_the_commands_actually_use(self): """A second definition of the autoroute dir must not drift from this one.""" @@ -365,7 +201,7 @@ class TestConflictingOwnersOfTheSettingsFile: class TestDoesNotDestroyUserOwnedStructure: - def test_writes_through_a_symlinked_settings_file(self, tmp_path, lite_on_path): + def test_writes_through_a_symlinked_settings_file(self, tmp_path): """os.replace() swaps the symlink for a regular file, detaching a dotfiles repo. There is no backup here to undo that, so the link must survive and its @@ -378,20 +214,20 @@ class TestDoesNotDestroyUserOwnedStructure: link.parent.mkdir() link.symlink_to(real) - _helper_configure("https://proxy.example.com", link, ()) + _static_configure("https://proxy.example.com", link, ()) assert link.is_symlink() assert json.loads(real.read_text())["env"]["ANTHROPIC_BASE_URL"] == "https://proxy.example.com" assert json.loads(real.read_text())["theme"] == "dark" - def test_refuses_rather_than_discarding_a_non_object_env(self, paths, lite_on_path): + def test_refuses_rather_than_discarding_a_non_object_env(self, paths): """merge coerces a non-dict env to {}; that is silent data loss on a persistent write.""" settings_path, backup_path = paths settings_path.parent.mkdir(parents=True) settings_path.write_text(json.dumps({"theme": "dark", "env": "not-an-object"})) with pytest.raises(ClaudeSettingsError, match="non-object"): - _helper_configure("https://proxy.example.com", settings_path, _owners(backup_path)) + _static_configure("https://proxy.example.com", settings_path, _owners(backup_path)) assert json.loads(settings_path.read_text())["env"] == "not-an-object" @@ -446,18 +282,13 @@ class TestConfigureStatePath: assert work_state == configure_state_path(tmp_path / "work" / "settings.json") def test_configure_and_unconfigure_under_a_config_dir_leave_the_default_receipt_alone( - self, default_paths, tmp_path, lite_on_path + self, default_paths, tmp_path ): _default_settings, default_state = default_paths work_settings = tmp_path / "work" / "settings.json" work_state = configure_state_path(work_settings) configure_claude_settings( - "https://proxy.example.com", - ApiKeyHelper(resolve_api_key_helper("https://proxy.example.com")), - KeepModel(), - work_settings, - work_state, - (), + "https://proxy.example.com", StaticToken("sk-virtual-key"), KeepModel(), work_settings, work_state, () ) assert work_state.exists() and not default_state.exists() outcome = unconfigure_claude_settings(work_settings, work_state, ()) @@ -465,51 +296,8 @@ class TestConfigureStatePath: assert not work_state.exists() -class TestLiteApiKeyHelperConfigured: - def _settings(self, tmp_path, payload): - settings_path = tmp_path / "settings.json" - settings_path.write_text(payload) - return settings_path - - def test_recognises_the_helper_lite_login_wrote_for_this_proxy(self, tmp_path, lite_on_path): - settings_path = tmp_path / "settings.json" - _helper_configure("https://proxy.example.com/", settings_path, (), tmp_path / "state.json") - - assert lite_api_key_helper_configured("https://proxy.example.com/", settings_path) is True - assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is True - - def test_a_helper_for_another_proxy_does_not_count(self, tmp_path, lite_on_path): - settings_path = tmp_path / "settings.json" - _helper_configure("https://other.example.com", settings_path, (), tmp_path / "state.json") - - assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False - - def test_a_hand_written_helper_does_not_count(self, tmp_path, lite_on_path): - settings_path = self._settings(tmp_path, json.dumps({"apiKeyHelper": "cat ~/.my-proxy-key"})) - - assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False - - def test_missing_or_helperless_settings_do_not_count(self, tmp_path, lite_on_path): - assert lite_api_key_helper_configured("https://proxy.example.com", tmp_path / "absent.json") is False - helperless = json.dumps({"env": {"ANTHROPIC_BASE_URL": "https://proxy.example.com"}}) - settings_path = self._settings(tmp_path, helperless) - assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False - - def test_unreadable_settings_fall_back_to_false(self, tmp_path, lite_on_path): - settings_path = self._settings(tmp_path, "{not json") - - assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False - - def test_lite_missing_from_path_falls_back_to_false(self, tmp_path): - helper = "/usr/local/bin/lite --base-url https://proxy.example.com auth print-token" - settings_path = self._settings(tmp_path, json.dumps({"apiKeyHelper": helper})) - with patch(f"{CLAUDE_SETTINGS_MODULE}.shutil.which", return_value=None): - assert lite_api_key_helper_configured("https://proxy.example.com", settings_path) is False - - class TestMergeClaudeSettings: - """One merge for every way Claude Code gets wired: `lite up`, `lite login --config-claude`, - `lite configure claude` and `lite autoroute up`.""" + """One merge for every way Claude Code gets wired: `lite up`, `lite configure claude` and `lite autoroute up`.""" def test_a_static_token_lands_in_env_and_the_helper_slot_is_cleared(self): settings = {"apiKeyHelper": "/usr/local/bin/lite auth print-token", "env": {"ANTHROPIC_API_KEY": "leaked"}} @@ -523,12 +311,6 @@ class TestMergeClaudeSettings: assert "model" not in merged assert not any(key in merged["env"] for key in ANTHROPIC_DEFAULT_MODEL_ENV_KEYS) - def test_a_helper_lands_top_level_and_the_static_slots_are_cleared(self): - settings = {"env": {"ANTHROPIC_AUTH_TOKEN": "sk-old", "ANTHROPIC_API_KEY": "leaked"}} - merged = merge_claude_settings(settings, "http://127.0.0.1:4000", ApiKeyHelper("lite auth print-token")) - assert merged["apiKeyHelper"] == "lite auth print-token" - assert "ANTHROPIC_AUTH_TOKEN" not in merged["env"] and "ANTHROPIC_API_KEY" not in merged["env"] - def test_keeps_existing_switch_values_and_unrelated_keys_without_mutating_the_input(self): settings = {"theme": "dark", "env": {"SOME_OTHER_VAR": "value", "ENABLE_TOOL_SEARCH": "false"}} merged = merge_claude_settings(settings, "http://127.0.0.1:4000", StaticToken("token-abc")) @@ -537,13 +319,22 @@ class TestMergeClaudeSettings: assert merged["env"]["ENABLE_TOOL_SEARCH"] == "false" assert settings == {"theme": "dark", "env": {"SOME_OTHER_VAR": "value", "ENABLE_TOOL_SEARCH": "false"}} - def test_a_default_model_sets_only_the_row_claude_code_starts_on(self): + def test_a_default_model_pins_the_starting_row_and_the_model_a_resumed_session_keeps(self): + # `model` is only the row Claude Code starts on: a resumed session re-sends the model its transcript + # recorded, which behind a raw-model auto-router is the tier model (403 for a key scoped to the + # router). ANTHROPIC_MODEL outranks the transcript on resume, so the pin has to land there too. merged = merge_claude_settings( {}, "http://127.0.0.1:4000", StaticToken("token-abc"), default_model="claude-auto" ) assert merged["model"] == "claude-auto" + assert merged["env"]["ANTHROPIC_MODEL"] == "claude-auto" assert not any(key in merged["env"] for key in ANTHROPIC_DEFAULT_MODEL_ENV_KEYS) + def test_without_a_default_model_neither_pin_is_written_and_a_users_own_stays(self): + settings = {"model": "mine", "env": {"ANTHROPIC_MODEL": "mine-too"}} + merged = merge_claude_settings(settings, "http://127.0.0.1:4000", StaticToken("token-abc")) + assert merged["model"] == "mine" and merged["env"]["ANTHROPIC_MODEL"] == "mine-too" + def test_a_tier_model_forces_every_claude_code_tier_as_autoroute_needs(self): # Router's auto-router registry is keyed by the literal requested model string with no # wildcard resolution, so `lite autoroute up` overrides the env var each tier reads. @@ -564,7 +355,7 @@ class TestMergeClaudeSettings: "apiKeyHelper": "old-helper", "model": "old-model", } - for credential in (StaticToken("token-abc"), ApiKeyHelper("helper")): + for credential in (StaticToken("token-abc"), StaticToken("token-rotated")): merged = merge_claude_settings(settings, "http://127.0.0.1:4000", credential, default_model="claude-auto") changed_top_level = {key for key in set(settings) | set(merged) if settings.get(key) != merged.get(key)} assert changed_top_level - {"env"} <= set(OWNED_TOP_LEVEL_KEYS) @@ -580,7 +371,7 @@ class TestMergeClaudeSettings: PROXY = "http://127.0.0.1:4000" ANTHROPIC = "https://api.anthropic.com" -HELPER = ApiKeyHelper("lite auth print-token") +RELOGIN = StaticToken("sk-fresh-login") ORIGINAL = { "theme": "dark", "permissions": {"allow": ["Bash"]}, @@ -664,18 +455,19 @@ UNDO_SCENARIOS = { { "restored": { "env.ANTHROPIC_BASE_URL", + "env.ANTHROPIC_AUTH_TOKEN", "env.ENABLE_TOOL_SEARCH", "env.CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY", - "apiKeyHelper", + "statusLine", }, "kept": (), }, - {"credential": HELPER, "model": KeepModel()}, + {"credential": RELOGIN, "model": KeepModel()}, ), "repeat across credential kinds keeps the first snapshot": ( ORIGINAL, [ - {"credential": HELPER, "model": UnpinModel()}, + {"credential": RELOGIN, "model": UnpinModel()}, {"credential": StaticToken("sk-rotated"), "model": StartOn("claude-sonnet-4-6")}, ], ORIGINAL, @@ -688,13 +480,13 @@ UNDO_SCENARIOS = { {"model": "claude-opus-5"}, {}, ), - "re-login keeps our pin": (None, [{"credential": HELPER, "model": KeepModel()}], None, {"file_removed": True}), + "re-login keeps our pin": (None, [{"credential": RELOGIN, "model": KeepModel()}], None, {"file_removed": True}), "edit between configures survives an unpin repeat": ( ORIGINAL, [ _set("model", "my-favourite"), _set("env.ENABLE_TOOL_SEARCH", "false"), - {"credential": HELPER, "model": UnpinModel()}, + {"credential": RELOGIN, "model": UnpinModel()}, ], {**ORIGINAL, "env": {**ORIGINAL["env"], "ENABLE_TOOL_SEARCH": "false"}, "model": "my-favourite"}, {"kept": {"env.ENABLE_TOOL_SEARCH", "model"}}, @@ -704,14 +496,14 @@ UNDO_SCENARIOS = { [ _set("model", "my-favourite"), _set("env.ENABLE_TOOL_SEARCH", "false"), - {"credential": HELPER, "model": KeepModel()}, + {"credential": RELOGIN, "model": KeepModel()}, ], {**ORIGINAL, "env": {**ORIGINAL["env"], "ENABLE_TOOL_SEARCH": "false"}, "model": "my-favourite"}, {"kept": {"env.ENABLE_TOOL_SEARCH", "model"}}, ), "edit between configures: a same-model repeat displaces it, so it is what comes back": ( ORIGINAL, - [_set("model", "my-favourite"), _set("env.ENABLE_TOOL_SEARCH", "false"), {"credential": HELPER}], + [_set("model", "my-favourite"), _set("env.ENABLE_TOOL_SEARCH", "false"), {"credential": RELOGIN}], {**ORIGINAL, "env": {**ORIGINAL["env"], "ENABLE_TOOL_SEARCH": "false"}, "model": "my-favourite"}, {"kept": {"env.ENABLE_TOOL_SEARCH"}, "restored_includes": {"model"}}, ), @@ -744,12 +536,12 @@ UNDO_SCENARIOS = { None, [ _set("env.ANTHROPIC_API_KEY", "sk-user"), - {"credential": HELPER, "model": KeepModel()}, + {"credential": RELOGIN, "model": KeepModel()}, _set("env.ANTHROPIC_BASE_URL", _ABSENT), ], None, {"withheld": {("env.ANTHROPIC_API_KEY", PROXY)}, "file_removed": True, "receipt_kept": True}, - {"credential": HELPER, "model": KeepModel()}, + {"credential": RELOGIN, "model": KeepModel()}, ), "a credential the user changed is kept, never also withheld": ( ORIGINAL, @@ -832,11 +624,10 @@ class TestConfigureAndUnconfigure: @pytest.mark.parametrize( ("path", "value", "repeat_credential"), [ - ("env.ANTHROPIC_API_KEY", "sk-user-added-later", HELPER), - ("env.ANTHROPIC_AUTH_TOKEN", "sk-users-own-token", HELPER), + ("env.ANTHROPIC_API_KEY", "sk-user-added-later", RELOGIN), ("apiKeyHelper", "/opt/mine/helper", StaticToken("sk-rotated")), ], - ids=["user-adds-api-key", "user-replaces-our-token", "user-sets-own-helper"], + ids=["user-adds-api-key", "user-sets-own-helper"], ) def test_a_credential_the_user_set_between_two_configures_is_what_comes_back( self, tmp_path, path, value, repeat_credential @@ -844,7 +635,7 @@ class TestConfigureAndUnconfigure: # The repeat's merge clears the slot, so the displaced value is snapshotted and is what returns; # it was set while the file pointed at the proxy, so it returns once the file points there again. rig = _Rig(tmp_path, {"theme": "dark"}) - rig.configure(credential=HELPER, model=KeepModel()) + rig.configure(credential=RELOGIN, model=KeepModel()) rig.edit(_set(path, value)) rig.configure(credential=repeat_credential, model=KeepModel()) assert not _lookup(rig.read(), path) @@ -955,3 +746,95 @@ class TestConfigureAndUnconfigure: def _lookup(settings, path): section, _, key = path.rpartition(".") return (settings.get(section) or {}).get(key) if section else settings.get(key) + + +class TestStatusLine: + """Every configure registers the status line, and only while the slot is empty or already ours.""" + + COMMAND = "/opt/lite/bin/python /Users/me/.litellm/statusline.py" + + def test_an_empty_slot_gets_our_status_line(self): + assert with_status_line({}, self.COMMAND)["statusLine"] == {"type": "command", "command": self.COMMAND} + + def test_a_users_own_status_line_is_never_replaced(self): + theirs = {"type": "command", "command": "~/.claude/my-statusline.sh"} + assert with_status_line({"statusLine": theirs}, self.COMMAND)["statusLine"] == theirs + + def test_ours_under_an_older_interpreter_is_refreshed(self): + stale = {"type": "command", "command": "/old/python /Users/me/.litellm/statusline.py"} + assert with_status_line({"statusLine": stale}, self.COMMAND)["statusLine"]["command"] == self.COMMAND + + def test_the_merge_carries_it(self): + merged = merge_claude_settings({}, PROXY, StaticToken("tok"), status_line=self.COMMAND) + assert merged["statusLine"] == {"type": "command", "command": self.COMMAND} + + def test_the_installed_script_is_the_bundled_one_and_the_command_runs_this_interpreter(self, tmp_path): + from litellm.proxy.client.cli.commands import statusline_script + + script = tmp_path / "lite" / "statusline.py" + command = install_statusline_script(script) + assert script.read_bytes() == pathlib.Path(statusline_script.__file__).read_bytes() + assert shlex.split(command) == [sys.executable, str(script)] + assert command == statusline_command(script) + assert stat.S_IMODE(script.stat().st_mode) == 0o600 + assert stat.S_IMODE(script.parent.stat().st_mode) == 0o700 + assert install_statusline_script(script) == command + + def test_a_reinstall_replaces_the_script_in_one_step_and_a_refused_one_leaves_the_old_script_whole(self, tmp_path): + # Claude Code may be running the script at the moment `lite` reinstalls it; the file it has open + # must stay complete, and a reinstall that cannot land must not leave a truncated script behind. + from litellm.proxy.client.cli.commands import statusline_script + + script = tmp_path / "lite" / "statusline.py" + install_statusline_script(script) + bundled = pathlib.Path(statusline_script.__file__).read_bytes() + with script.open("rb") as running: + install_statusline_script(script) + assert running.read() == bundled + assert [child.name for child in script.parent.iterdir()] == ["statusline.py"] + + if os.geteuid() != 0: + script.parent.chmod(0o500) + try: + with pytest.raises(ClaudeSettingsError, match="Could not install the status line script"): + install_statusline_script(script) + finally: + script.parent.chmod(0o700) + assert script.read_bytes() == bundled + + def test_configure_installs_it_and_unconfigure_removes_only_ours(self, tmp_path): + rig = _Rig(tmp_path, {"theme": "dark"}) + script = tmp_path / "statusline.py" + rig.configure(script_path=script) + assert rig.read()["statusLine"]["command"] == statusline_command(script) + assert script.exists() + + outcome = rig.unconfigure() + assert rig.read() == {"theme": "dark"} + assert "statusLine" in outcome.restored + + def test_a_status_line_the_user_replaced_after_configure_survives_unconfigure(self, tmp_path): + rig = _Rig(tmp_path, None) + rig.configure(model=KeepModel(), script_path=tmp_path / "statusline.py") + theirs = {"type": "command", "command": "~/.claude/my-statusline.sh"} + rig.edit(lambda settings: {**settings, "statusLine": theirs}) + + outcome = rig.unconfigure() + assert rig.read()["statusLine"] == theirs + assert "statusLine" in outcome.kept + + def test_a_receipt_from_before_the_status_line_existed_still_unconfigures(self, tmp_path): + # Older receipts never claimed statusLine; a key no configure wrote is never ours, so it stays. + rig = _Rig(tmp_path, None) + script = tmp_path / "statusline.py" + rig.configure(model=KeepModel(), script_path=script) + receipt = json.loads(rig.state.read_text()) + receipt["written"].pop("statusLine") + receipt["previous"].pop("statusLine") + rig.state.write_text(json.dumps(receipt)) + + outcome = rig.unconfigure() + restored = rig.read() + assert "ANTHROPIC_AUTH_TOKEN" not in restored.get("env", {}) + assert restored["statusLine"]["command"] == statusline_command(script) + assert "statusLine" not in outcome.restored diff --git a/tests/test_litellm/proxy/client/cli/test_configure_commands.py b/tests/test_litellm/proxy/client/cli/test_configure_commands.py index ac7408339fe..8ed188af737 100644 --- a/tests/test_litellm/proxy/client/cli/test_configure_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_configure_commands.py @@ -86,6 +86,7 @@ class TestConfigureClaudeWithAVirtualKey: assert "ANTHROPIC_DEFAULT_SONNET_MODEL" not in written["env"] assert state_path.exists() assert VALID_KEY not in result.output + assert written["env"]["ANTHROPIC_MODEL"] == "claude-auto" assert "Starting model: claude-auto" in result.output assert "1 of the proxy's 2 models" in result.output assert "lite unconfigure claude" in result.output @@ -99,7 +100,7 @@ class TestConfigureClaudeWithAVirtualKey: assert result.exit_code == 0, result.output written = json.loads(settings_path.read_text()) assert written["env"]["ANTHROPIC_AUTH_TOKEN"] == VALID_KEY - assert "model" not in written + assert "model" not in written and "ANTHROPIC_MODEL" not in written["env"] assert "Starting model: not pinned" in result.output @responses.activate @@ -159,18 +160,11 @@ class TestConfigureClaudeWithAVirtualKey: assert not settings_path.exists() @responses.activate - @pytest.mark.parametrize("entry", ["virtual-key", "login", "interactive"]) - def test_refuses_while_lite_up_holds_a_backup_before_any_login_or_request( - self, runner, paths, monkeypatch, lite_up_backup, entry - ): + @pytest.mark.parametrize("entry", ["virtual-key", "no-key", "interactive"]) + def test_refuses_while_lite_up_holds_a_backup_before_any_request(self, runner, paths, lite_up_backup, entry): _mock_models() - - def login_must_not_run(ctx): - raise AssertionError("the local precondition must be checked before a login is attempted") - - monkeypatch.setattr(configure_module, "ensure_fresh_login", login_must_not_run) if entry == "interactive": - ctx = click.Context(configure_group, obj={"base_url": PROXY, "api_key": None}) + ctx = click.Context(configure_group, obj={"base_url": PROXY, "api_key": VALID_KEY}) with pytest.raises(click.ClickException, match="lite down"): interactive_configure(ctx, pick_targets=lambda: ("claude",), pick_model=lambda listed: None) else: @@ -195,33 +189,27 @@ class TestConfigureClaudeWithAVirtualKey: assert json.loads(target.read_text())["env"]["ANTHROPIC_AUTH_TOKEN"] == VALID_KEY -class TestConfigureClaudeWithTheLogin: - def _stored_login(self, monkeypatch): - monkeypatch.setattr(configure_module, "ensure_fresh_login", lambda ctx: None) - monkeypatch.setattr(configure_module, "get_stored_api_key", lambda expected_base_url, vault: VALID_KEY) - +class TestConfigureClaudeWithoutAKey: @responses.activate - def test_uses_the_login_through_the_helper_and_writes_no_secret(self, runner, paths, monkeypatch, lite_on_path): + def test_refuses_and_names_the_ways_to_pass_a_key_without_writing_or_logging_in(self, runner, paths): + # A `lite login` credential expires within a day; the old fallback wrote an apiKeyHelper that made + # Claude Code spawn `lite` (and its keychain probe) on every credential refresh. _mock_models() - self._stored_login(monkeypatch) - settings_path, _ = paths + settings_path, state_path = paths result = runner.invoke( configure_claude, ["--model", "claude-auto"], - obj={"base_url": PROXY, "api_key": VALID_KEY, "api_key_from_token_file": True}, + obj={"base_url": PROXY, "api_key": "sk-login-jwt", "api_key_from_token_file": True}, ) - assert result.exit_code == 0, result.output - written = json.loads(settings_path.read_text()) - assert written["apiKeyHelper"] == f"{lite_on_path} --base-url {PROXY} auth print-token" - assert "ANTHROPIC_AUTH_TOKEN" not in written["env"] - assert written["model"] == "claude-auto" - assert VALID_KEY not in settings_path.read_text() - assert "read through apiKeyHelper" in result.output + assert result.exit_code != 0 + assert "--api-key" in result.output and "LITELLM_PROXY_API_KEY" in result.output + assert "apiKeyHelper" not in result.output + assert not settings_path.exists() and not state_path.exists() + assert len(responses.calls) == 0 @responses.activate - def test_an_explicit_key_still_wins_over_a_stored_login(self, runner, paths, monkeypatch, lite_on_path): + def test_an_explicit_key_still_wins_over_a_stored_login(self, runner, paths): _mock_models() - self._stored_login(monkeypatch) settings_path, _ = paths result = runner.invoke( configure_claude, @@ -302,11 +290,13 @@ class TestUnconfigureClaude: edited = json.loads(settings_path.read_text()) edited["env"] = {key: f"{value}-edited" for key, value in edited["env"].items()} edited["model"] = "mine" + edited["statusLine"] = {"type": "command", "command": "~/.claude/my-statusline.sh"} settings_path.write_text(json.dumps(edited)) result = runner.invoke(cli, ["unconfigure", "claude"]) assert result.exit_code == 0, result.output assert "Nothing in" in result.output and "was still ours to restore" in result.output assert "Left as you changed them since:" in result.output and "model" in result.output + assert "statusLine" in result.output @responses.activate def test_names_the_server_a_withheld_credential_was_captured_with_and_keeps_the_receipt(self, runner, paths): diff --git a/tests/test_litellm/proxy/client/cli/test_statusline_script.py b/tests/test_litellm/proxy/client/cli/test_statusline_script.py new file mode 100644 index 00000000000..691764fbef4 --- /dev/null +++ b/tests/test_litellm/proxy/client/cli/test_statusline_script.py @@ -0,0 +1,404 @@ +"""The status line script is copied verbatim to the user's machine, so these drive it the way Claude Code +and Codex do: the documented stdin payload, a transcript on disk, and the proxy behind an injected fetch.""" + +import io +import json +import os +import re +import subprocess +import sys +from pathlib import Path + +import pytest + +from litellm.proxy.client.cli.commands import statusline_script +from litellm.proxy.client.cli.commands.statusline_script import ( + CACHE_TTL_SECONDS, + Credentials, + Fetched, + Session, + cache_dir_name, + cache_path, + claude_credentials, + codex_credentials, + latest_transcript_model, + load_session, + render, + run, +) + +SESSION_ID = "cf712ab8-4c7c-4d48-ba91-eed54bc2956b" +ANSI = re.compile(r"\x1b\[[0-9;]*m") +RECORDED = Session( + router_name="claude-auto", + last_model="anthropic/claude-sonnet-5", + spend=0.14, + baseline_spend=0.38, + baseline_model="anthropic/claude-opus-5", +) + + +def _assistant_line(model: str, **extra: object) -> str: + return json.dumps({"type": "assistant", "message": {"model": model, "role": "assistant"}, **extra}) + + +@pytest.fixture +def transcript(tmp_path: Path) -> Path: + path = tmp_path / "session.jsonl" + path.write_text( + "\n".join( + ( + json.dumps({"type": "user", "message": {"role": "user", "content": "hi"}}), + _assistant_line("claude-haiku-4-5"), + json.dumps({"type": "user", "message": {"role": "user", "content": "harder"}}), + _assistant_line("claude-sonnet-5"), + _assistant_line("claude-haiku-4-5", isSidechain=True), + _assistant_line("claude-haiku-4-5", agentId="agent-1"), + json.dumps({"type": "progress", "data": {}}), + ) + ) + + "\n" + ) + return path + + +@pytest.fixture +def config_dir(tmp_path: Path) -> Path: + directory = tmp_path / "claude" + (directory / "cache").mkdir(parents=True) + (directory / "cache" / "gateway-models.json").write_text( + json.dumps({"models": [{"id": "claude-opus-5", "display_name": "Claude Opus 5"}]}) + ) + return directory + + +def _payload(transcript: Path, session_id: str = SESSION_ID) -> dict: + return { + "session_id": session_id, + "transcript_path": str(transcript), + "model": {"id": "claude-auto", "display_name": "claude-auto"}, + } + + +def _env(tmp_path: Path, config_dir: Path, **extra: str) -> dict[str, str]: + return { + "TMPDIR": str(tmp_path / "tmp"), + "CLAUDE_CONFIG_DIR": str(config_dir), + "TERM": "dumb", + "ANTHROPIC_BASE_URL": "http://127.0.0.1:4000", + "ANTHROPIC_AUTH_TOKEN": "sk-virtual", + **extra, + } + + +def _run(payload: object, env: dict[str, str], fetch) -> str: + out = io.StringIO() + run(io.StringIO(json.dumps(payload)), out, env, fetch) + return out.getvalue() + + +class TestTranscript: + def test_the_latest_foreground_assistant_line_wins_over_later_sidechain_and_agent_lines(self, transcript): + assert latest_transcript_model(str(transcript)) == "claude-sonnet-5" + + def test_a_synthetic_line_is_not_a_served_model(self, tmp_path): + # Claude Code writes `` for messages it produced locally (an API error on resume, for one); + # showing "Routed to: " would name a model no proxy served. + path = tmp_path / "t.jsonl" + path.write_text(_assistant_line("claude-haiku-4-5") + "\n" + _assistant_line("") + "\n") + assert latest_transcript_model(str(path)) == "claude-haiku-4-5" + + def test_a_missing_or_empty_transcript_yields_nothing(self, tmp_path): + empty = tmp_path / "empty.jsonl" + empty.write_text("") + assert latest_transcript_model(str(tmp_path / "missing.jsonl")) == "" + assert latest_transcript_model(str(empty)) == "" + assert latest_transcript_model("") == "" + + +class TestCredentials: + MIXED = { + "ANTHROPIC_BASE_URL": "http://anthropic-side:4000", + "ANTHROPIC_AUTH_TOKEN": "sk-ant", + "OPENAI_BASE_URL": "http://openai-side:4000/v1/", + "OPENAI_API_KEY": "sk-openai", + } + + def test_each_agent_reads_the_pair_it_dials_itself(self): + # A shell that exports both families must not send Codex's hook to the Anthropic proxy. + assert claude_credentials(self.MIXED) == Credentials("http://anthropic-side:4000", "sk-ant") + assert codex_credentials(self.MIXED) == Credentials("http://openai-side:4000", "sk-openai") + + def test_lites_own_shell_variables_are_not_a_credential_either_agent_sends(self): + # A `lite login` shell exports LITELLM_PROXY_*; Claude Code and Codex never read them, so the + # status line must not query the proxy as that principal while the agent used another. + env = {"LITELLM_PROXY_URL": "http://lite:4000/", "LITELLM_PROXY_API_KEY": "sk-lite", **self.MIXED} + assert claude_credentials(env) == Credentials("http://anthropic-side:4000", "sk-ant") + assert codex_credentials(env) == Credentials("http://openai-side:4000", "sk-openai") + assert not claude_credentials({"LITELLM_PROXY_API_KEY": "sk-lite", "ANTHROPIC_BASE_URL": "http://p"}).usable + assert codex_credentials({}) == Credentials("", "") + + def test_claude_code_prefers_the_auth_token_over_a_stray_api_key(self): + env = {"ANTHROPIC_BASE_URL": "http://p", "ANTHROPIC_API_KEY": "sk-stray", "ANTHROPIC_AUTH_TOKEN": "sk-ours"} + assert claude_credentials(env).api_key == "sk-ours" + + def test_an_api_key_helper_in_settings_is_never_run(self, tmp_path, transcript, config_dir): + # `lite` once wrote `apiKeyHelper: lite auth print-token`; running it from a status line that + # refreshes every 300ms spawned `lite` (and a keychain prompt) on every refresh. Without a key + # in the env the proxy is simply not asked. + (config_dir / "settings.json").write_text(json.dumps({"apiKeyHelper": "printf sk-from-helper"})) + asked = [] + env = {k: v for k, v in _env(tmp_path, config_dir).items() if k != "ANTHROPIC_AUTH_TOKEN"} + text = _run(_payload(transcript), env, lambda c, s: asked.append(c) or Fetched(RECORDED, True)) + assert text == "Routed to: claude-sonnet-5" and asked == [] + + +class TestSessionCache: + def test_a_definite_answer_is_served_from_the_cache_within_the_ttl(self, tmp_path): + calls = [] + + def fetch(credentials, session_id): + calls.append(session_id) + return Fetched(RECORDED, definitive=True) + + clock = [100.0] + credentials = Credentials("http://p", "sk") + first = load_session(credentials, SESSION_ID, tmp_path, fetch, now=lambda: clock[0]) + clock[0] = 100.0 + CACHE_TTL_SECONDS - 1 + second = load_session(credentials, SESSION_ID, tmp_path, fetch, now=lambda: clock[0]) + clock[0] = 100.0 + CACHE_TTL_SECONDS + 1 + third = load_session(credentials, SESSION_ID, tmp_path, fetch, now=lambda: clock[0]) + assert first == second == third == RECORDED + assert calls == [SESSION_ID, SESSION_ID] + + def test_a_404_is_cached_as_absence_but_a_transport_failure_is_retried(self, tmp_path): + outcomes = iter((Fetched(None, definitive=False), Fetched(None, definitive=True), Fetched(RECORDED, True))) + calls = [] + + def fetch(credentials, session_id): + calls.append(session_id) + return next(outcomes) + + credentials = Credentials("http://p", "sk") + assert load_session(credentials, SESSION_ID, tmp_path, fetch, now=lambda: 1.0) is None + assert load_session(credentials, SESSION_ID, tmp_path, fetch, now=lambda: 1.0) is None + assert load_session(credentials, SESSION_ID, tmp_path, fetch, now=lambda: 1.0) is None + assert len(calls) == 2 + + def test_a_cache_directory_that_is_not_private_is_never_used(self, tmp_path): + # A shared temp root lets another user pre-create the directory; refuse it rather than write into it. + calls = [] + + def fetch(credentials, session_id): + calls.append(session_id) + return Fetched(RECORDED, definitive=True) + + shared = tmp_path / "litellm-statusline" + shared.mkdir(mode=0o755) + credentials = Credentials("http://p", "sk") + for _ in range(2): + assert load_session(credentials, SESSION_ID, shared, fetch, now=lambda: 1.0) == RECORDED + assert calls == [SESSION_ID, SESSION_ID] + assert list(shared.iterdir()) == [] + + def test_any_client_error_is_a_definite_answer_and_a_server_error_is_not(self, monkeypatch): + import urllib.error + + from litellm.proxy.client.cli.commands.statusline_script import fetch_session + + def fail_with(code): + def opener(request, timeout): + raise urllib.error.HTTPError(request.full_url, code, "x", {}, None) + + return opener + + for code, definitive in ((403, True), (401, True), (404, True), (502, False)): + monkeypatch.setattr("urllib.request.urlopen", fail_with(code)) + assert fetch_session(Credentials("http://127.0.0.1:1", "sk"), SESSION_ID) == Fetched(None, definitive) + + def test_the_cache_file_holds_the_proxy_answer_and_never_the_key(self, tmp_path): + credentials = Credentials("http://p", "sk-secret") + load_session(credentials, SESSION_ID, tmp_path, lambda c, s: Fetched(RECORDED, True)) + path = cache_path(tmp_path, credentials, SESSION_ID) + assert "sk-secret" not in written and SESSION_ID not in written if (written := path.read_text()) else False + assert "sk-secret" not in path.name + assert json.loads(written)["session"]["baseline_model"] == "anthropic/claude-opus-5" + assert (path.stat().st_mode & 0o777) == 0o600 + assert (path.parent.stat().st_mode & 0o777) == 0o700 + + def test_a_refresh_replaces_the_entry_in_one_step_so_a_concurrent_refresh_never_reads_a_torn_one(self, tmp_path): + credentials = Credentials("http://p", "sk") + load_session(credentials, SESSION_ID, tmp_path, lambda c, s: Fetched(RECORDED, True), now=lambda: 1.0) + path = cache_path(tmp_path, credentials, SESSION_ID) + first = path.read_text() + + with path.open() as concurrent_reader: + newer = RECORDED._replace(spend=0.5) + load_session(credentials, SESSION_ID, tmp_path, lambda c, s: Fetched(newer, True), now=lambda: 100.0) + assert concurrent_reader.read() == first + assert json.loads(path.read_text())["session"]["spend"] == 0.5 + assert (path.stat().st_mode & 0o777) == 0o600 + assert [child.name for child in tmp_path.iterdir()] == [path.name] + + def test_the_same_session_id_against_another_proxy_or_key_is_not_served_from_the_cache(self, tmp_path): + answers = iter((Fetched(RECORDED, True), Fetched(RECORDED._replace(spend=9.0), True))) + first = load_session(Credentials("http://p", "sk-a"), SESSION_ID, tmp_path, lambda c, s: next(answers)) + second = load_session(Credentials("http://p", "sk-b"), SESSION_ID, tmp_path, lambda c, s: next(answers)) + assert first == RECORDED and second is not None and second.spend == 9.0 + + +class TestRender: + def test_savings_header_and_bars_against_the_routers_baseline(self, config_dir): + text = render("claude-sonnet-5", RECORDED, config_dir, use_color=False, bar_width=10) + assert text.splitlines() == [ + "claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5", + "LiteLLM ████░░░░░░ $0.14", + "Claude Opus 5 ██████████ $0.38", + ] + + def test_control_characters_in_any_externally_sourced_label_never_reach_the_terminal(self, tmp_path, config_dir): + # The transcript, the proxy payload and Claude Code's model cache all feed labels straight into a + # terminal, and none is under this script's control. Only the control bytes are dropped (ESC, BEL, + # C1), which is what disarms an OSC-52 clipboard write or a screen clear; the printable remainder of + # such a sequence is inert text and is kept as is. + from litellm.proxy.client.cli.commands.statusline_script import _session_from_payload, model_label + + hostile = "claude-\x1b\x07\x9bsonnet" + path = tmp_path / "t.jsonl" + path.write_text(_assistant_line(hostile) + "\n") + assert latest_transcript_model(str(path)) == "claude-sonnet" + (config_dir / "cache" / "gateway-models.json").write_text( + json.dumps({"models": [{"id": "claude-sonnet", "display_name": "Son\x1b\x07net"}]}) + ) + assert model_label("claude-sonnet", config_dir) == "Sonnet" + session = _session_from_payload( + {"router_name": "auto\x07", "last_model": hostile, "spend": 0.1, "baseline_spend": 0.2, "baseline_model": "op\x1bus"} + ) + assert session == Session("auto", "claude-sonnet", 0.1, 0.2, "opus") + assert latest_transcript_model(str(path)) == "claude-sonnet" + assert "\x1b]52;c;ZXZpbA==" not in render( + latest_transcript_model(str(path)), + _session_from_payload({"router_name": "a", "last_model": "m", "spend": 0.1, "baseline_spend": 0.2, "baseline_model": "\x1b]52;c;ZXZpbA==\x07"}), + config_dir, + use_color=False, + ) + + def test_a_session_that_cost_more_than_its_baseline_reads_as_a_plus(self, config_dir): + dearer = RECORDED._replace(spend=0.50, baseline_spend=0.40) + assert "+25% vs Claude Opus 5" in render("m", dearer, config_dir, use_color=False) + + def test_without_a_baseline_only_the_routed_line_shows(self, config_dir): + assert render("m", RECORDED._replace(baseline_model=None), config_dir, False) == "claude-auto · Routed to: m" + assert render("m", None, config_dir, False) == "Routed to: m" + + def test_color_wraps_the_same_text(self, config_dir): + colored = render("claude-sonnet-5", RECORDED, config_dir, use_color=True, bar_width=10) + assert ANSI.sub("", colored) == render("claude-sonnet-5", RECORDED, config_dir, use_color=False, bar_width=10) + + +class TestClaudeCodeMode: + def test_the_transcript_names_the_routed_model_and_the_proxy_adds_the_savings(self, tmp_path, transcript, config_dir): + seen = [] + + def fetch(credentials, session_id): + seen.append((credentials, session_id)) + return Fetched(RECORDED, definitive=True) + + text = _run(_payload(transcript), _env(tmp_path, config_dir), fetch) + assert text.startswith("claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5\n") + assert seen == [(Credentials("http://127.0.0.1:4000", "sk-virtual"), SESSION_ID)] + + def test_an_unrecorded_session_degrades_to_the_routed_line(self, tmp_path, transcript, config_dir): + assert _run(_payload(transcript), _env(tmp_path, config_dir), lambda c, s: Fetched(None, True)) == ( + "Routed to: claude-sonnet-5" + ) + + def test_without_credentials_the_proxy_is_never_asked(self, tmp_path, transcript, config_dir): + def fetch(credentials, session_id): + raise AssertionError("must not fetch") + + env = {k: v for k, v in _env(tmp_path, config_dir).items() if k != "ANTHROPIC_AUTH_TOKEN"} + assert _run(_payload(transcript), env, fetch) == "Routed to: claude-sonnet-5" + + def test_before_the_first_response_the_payloads_display_name_shows(self, tmp_path, config_dir): + payload = _payload(tmp_path / "missing.jsonl") + assert _run(payload, _env(tmp_path, config_dir), lambda c, s: Fetched(RECORDED, True)) == "claude-auto" + + def test_a_discovered_display_name_labels_the_routed_model(self, tmp_path, config_dir): + path = tmp_path / "t.jsonl" + path.write_text(_assistant_line("anthropic/claude-opus-5") + "\n") + assert _run(_payload(path), _env(tmp_path, config_dir), lambda c, s: Fetched(None, True)) == ( + "Routed to: Claude Opus 5" + ) + + def test_the_cache_lands_under_the_platforms_temp_dir(self, tmp_path, transcript, config_dir): + env = {k: v for k, v in _env(tmp_path, config_dir).items() if k != "TMPDIR"} + env["TEMP"] = str(tmp_path / "wintemp") + _run(_payload(transcript), env, lambda c, s: Fetched(RECORDED, True)) + assert (tmp_path / "wintemp" / cache_dir_name()).is_dir() + assert cache_dir_name().endswith(str(os.getuid())) + + def test_a_crash_falls_back_to_the_model_label_claude_code_already_knows(self, tmp_path, transcript, config_dir): + def fetch(credentials, session_id): + raise RuntimeError("boom") + + assert _run(_payload(transcript), _env(tmp_path, config_dir), fetch) == "claude-auto" + + def test_garbage_on_stdin_still_prints_something(self, tmp_path, config_dir): + out = io.StringIO() + run(io.StringIO("not json"), out, _env(tmp_path, config_dir), lambda c, s: Fetched(None, True)) + assert out.getvalue() == "claude" + + +class TestCodexMode: + def test_the_stop_hook_prints_a_system_message_from_the_proxys_record(self, tmp_path, config_dir): + env = _env(tmp_path, config_dir, OPENAI_BASE_URL="http://127.0.0.1:4000/v1", OPENAI_API_KEY="sk-codex") + env = {k: v for k, v in env.items() if not k.startswith("ANTHROPIC_")} + seen = [] + + def fetch(credentials, session_id): + seen.append(credentials) + return Fetched(RECORDED, definitive=True) + + out = _run({"hook_event_name": "Stop", "session_id": SESSION_ID, "transcript_path": "/nope"}, env, fetch) + message = json.loads(out)["systemMessage"] + assert message.splitlines()[1] == "claude-auto · Routed to: claude-sonnet-5 -63% vs Claude Opus 5" + assert message.startswith("\n") + assert seen == [Credentials("http://127.0.0.1:4000", "sk-codex")] + + def test_an_unrecorded_session_prints_nothing_so_codex_shows_no_message(self, tmp_path, config_dir): + payload = {"hook_event_name": "Stop", "session_id": SESSION_ID} + assert _run(payload, _env(tmp_path, config_dir), lambda c, s: Fetched(None, True)) == "" + + def test_a_crash_prints_nothing_rather_than_text_codex_would_reject(self, tmp_path, config_dir): + def fetch(credentials, session_id): + raise RuntimeError("boom") + + env = _env(tmp_path, config_dir, OPENAI_BASE_URL="http://127.0.0.1:4000/v1", OPENAI_API_KEY="sk-codex") + assert _run({"hook_event_name": "Stop", "session_id": SESSION_ID}, env, fetch) == "" + + def test_a_turn_right_after_an_unrecorded_one_still_asks_the_proxy(self, tmp_path, config_dir): + # One hook run per turn: a miss on turn one must not be cached across turn two's fetch. + answers = iter((Fetched(None, definitive=True), Fetched(RECORDED, definitive=True))) + payload = {"hook_event_name": "Stop", "session_id": SESSION_ID} + env = _env(tmp_path, config_dir, OPENAI_BASE_URL="http://127.0.0.1:4000/v1", OPENAI_API_KEY="sk-codex") + assert _run(payload, env, lambda c, s: next(answers)) == "" + assert "Routed to: claude-sonnet-5" in json.loads(_run(payload, env, lambda c, s: next(answers)))["systemMessage"] + + +class TestStandalone: + def test_the_file_runs_under_a_bare_interpreter_with_no_litellm_on_the_path(self, tmp_path, transcript, config_dir): + # It is copied verbatim to ~/.litellm/statusline.py, so it must be self-contained. + script = tmp_path / "statusline.py" + script.write_bytes(Path(statusline_script.__file__).read_bytes()) + env = {k: v for k, v in _env(tmp_path, config_dir).items() if k != "ANTHROPIC_AUTH_TOKEN"} + completed = subprocess.run( + [sys.executable, "-I", str(script)], + input=json.dumps(_payload(transcript)), + capture_output=True, + text=True, + env=env, + check=True, + timeout=30, + ) + assert completed.stdout == "Routed to: claude-sonnet-5" diff --git a/tests/test_litellm/proxy/client/cli/test_up_commands.py b/tests/test_litellm/proxy/client/cli/test_up_commands.py index 111d2f3d682..ddd54dd1374 100644 --- a/tests/test_litellm/proxy/client/cli/test_up_commands.py +++ b/tests/test_litellm/proxy/client/cli/test_up_commands.py @@ -11,7 +11,7 @@ from click.testing import CliRunner from litellm.proxy.client.cli.commands import up as up_module from litellm.proxy.client.cli.commands.agents import AgentRunError -from litellm.proxy.client.cli.commands.claude_settings import ApiKeyHelper, ClaudeSettingsError +from litellm.proxy.client.cli.commands.claude_settings import ClaudeSettingsError, StaticToken from litellm.proxy.client.cli.commands.up import ( BackupRecord, UpError, @@ -20,7 +20,6 @@ from litellm.proxy.client.cli.commands.up import ( load_json_or_empty, merge_claude_settings, read_backup, - resolve_api_key_helper, restore_claude_settings, up, write_backup, @@ -40,52 +39,54 @@ def _patch_paths(monkeypatch, tmp_path): class TestMergeClaudeSettings: def test_preserves_unrelated_top_level_keys(self): - merged = merge_claude_settings({"theme": "dark"}, "http://localhost:4000", ApiKeyHelper("helper")) + merged = merge_claude_settings({"theme": "dark"}, "http://localhost:4000", StaticToken("sk-fresh"), status_line="statusline-cmd") assert merged["theme"] == "dark" def test_preserves_unrelated_env_keys(self): settings = {"env": {"SOME_OTHER_VAR": "value"}} - merged = merge_claude_settings(settings, "http://localhost:4000", ApiKeyHelper("helper")) + merged = merge_claude_settings(settings, "http://localhost:4000", StaticToken("sk-fresh"), status_line="statusline-cmd") assert merged["env"]["SOME_OTHER_VAR"] == "value" - def test_overrides_base_url_and_helper(self): + def test_overrides_base_url_and_strips_an_old_helper(self): settings = { "env": {"ANTHROPIC_BASE_URL": "https://old.example.com"}, "apiKeyHelper": "old-helper", } - merged = merge_claude_settings(settings, "http://localhost:4000/", ApiKeyHelper("new-helper")) + merged = merge_claude_settings(settings, "http://localhost:4000/", StaticToken("sk-fresh"), status_line="statusline-cmd") assert merged["env"]["ANTHROPIC_BASE_URL"] == "http://localhost:4000" + assert merged["env"]["ANTHROPIC_AUTH_TOKEN"] == "sk-fresh" assert merged["env"]["ENABLE_TOOL_SEARCH"] == "true" assert merged["env"]["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "1" - assert merged["apiKeyHelper"] == "new-helper" + assert "apiKeyHelper" not in merged def test_preserves_existing_gateway_model_discovery(self): settings = {"env": {"CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "0"}} - merged = merge_claude_settings(settings, "http://localhost:4000", ApiKeyHelper("helper")) + merged = merge_claude_settings(settings, "http://localhost:4000", StaticToken("sk-fresh"), status_line="statusline-cmd") assert merged["env"]["CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY"] == "0" def test_preserves_existing_tool_search(self): settings = {"env": {"ENABLE_TOOL_SEARCH": "false"}} - merged = merge_claude_settings(settings, "http://localhost:4000", ApiKeyHelper("helper")) + merged = merge_claude_settings(settings, "http://localhost:4000", StaticToken("sk-fresh"), status_line="statusline-cmd") assert merged["env"]["ENABLE_TOOL_SEARCH"] == "false" def test_drops_stray_api_key(self): settings = {"env": {"ANTHROPIC_API_KEY": "leaked-key"}} - merged = merge_claude_settings(settings, "http://localhost:4000", ApiKeyHelper("helper")) + merged = merge_claude_settings(settings, "http://localhost:4000", StaticToken("sk-fresh"), status_line="statusline-cmd") assert "ANTHROPIC_API_KEY" not in merged["env"] def test_works_from_empty_settings(self): - merged = merge_claude_settings({}, "http://localhost:4000", ApiKeyHelper("helper")) + merged = merge_claude_settings({}, "http://localhost:4000", StaticToken("sk-fresh"), status_line="statusline-cmd") assert merged["env"] == { "ANTHROPIC_BASE_URL": "http://localhost:4000", + "ANTHROPIC_AUTH_TOKEN": "sk-fresh", "ENABLE_TOOL_SEARCH": "true", "CLAUDE_CODE_ENABLE_GATEWAY_MODEL_DISCOVERY": "1", } - assert merged["apiKeyHelper"] == "helper" + assert "apiKeyHelper" not in merged def test_does_not_mutate_input(self): settings = {"env": {"FOO": "bar"}} - merge_claude_settings(settings, "http://localhost:4000", ApiKeyHelper("helper")) + merge_claude_settings(settings, "http://localhost:4000", StaticToken("sk-fresh"), status_line="statusline-cmd") assert settings == {"env": {"FOO": "bar"}} @@ -154,6 +155,25 @@ class TestBackupRoundTrip: assert restore_claude_settings() is None assert not settings_path.exists() + def test_restore_writes_owner_only_and_through_a_symlink(self, monkeypatch, tmp_path): + # The backup can hold a token the user had in the file before `up`; a plain open() would put it + # back under the umask, and would replace a dotfiles symlink with a regular file. + target = tmp_path / "dotfiles" / "settings.json" + target.parent.mkdir() + target.write_text("{}") + target.chmod(0o644) + settings_path = tmp_path / "settings.json" + settings_path.symlink_to(target) + monkeypatch.setattr(up_module, "CLAUDE_SETTINGS_PATH", settings_path) + monkeypatch.setattr(up_module, "BACKUP_PATH", tmp_path / "backup.json") + write_backup(BackupRecord(existed=True, content={"env": {"ANTHROPIC_AUTH_TOKEN": "sk-theirs"}})) + + restore_claude_settings() + + assert settings_path.is_symlink() + assert json.loads(target.read_text()) == {"env": {"ANTHROPIC_AUTH_TOKEN": "sk-theirs"}} + assert stat.S_IMODE(target.stat().st_mode) == 0o600 + def test_recreates_claude_dir_if_it_was_deleted_while_up_was_running(self, monkeypatch, tmp_path): """If ~/.claude/ is removed while `lite up` holds it open, restoring must recreate the directory rather than crash with FileNotFoundError and strand the backup file, which @@ -216,49 +236,6 @@ class TestBackupRoundTrip: assert not backup_path.exists() -class TestResolveApiKeyHelper: - def test_returns_helper_command_bound_to_the_selected_proxy(self, monkeypatch): - monkeypatch.setattr(shutil, "which", lambda name: "/usr/local/bin/lite") - helper = resolve_api_key_helper("http://localhost:4000") - assert helper == "/usr/local/bin/lite --base-url http://localhost:4000 auth print-token" - - def test_quotes_a_base_url_containing_shell_metacharacters(self, monkeypatch): - monkeypatch.setattr(shutil, "which", lambda name: "/usr/local/bin/lite") - helper = resolve_api_key_helper("http://example.com/path; rm -rf /") - assert helper == "/usr/local/bin/lite --base-url 'http://example.com/path; rm -rf /' auth print-token" - - def test_raises_when_lite_not_on_path(self, monkeypatch): - monkeypatch.setattr(shutil, "which", lambda name: None) - with pytest.raises(ClaudeSettingsError, match="Could not find `lite`"): - resolve_api_key_helper("http://localhost:4000") - - def test_windows_quotes_for_cmd_exe_instead_of_posix_sh(self, monkeypatch): - """cmd.exe takes a single quote literally, so a POSIX-quoted backslashed path is unrunnable.""" - lite_exe = "C:\\Users\\u\\AppData\\Local\\Programs\\Python\\Python313\\Scripts\\lite.EXE" - monkeypatch.setattr(shutil, "which", lambda name: lite_exe) - - helper = resolve_api_key_helper("https://gateway.example.com", platform="win32") - - assert helper == f'"{lite_exe}" "--base-url" "https://gateway.example.com" "auth" "print-token"' - - def test_windows_keeps_a_spaced_path_and_a_metacharacter_url_as_single_tokens(self, monkeypatch): - monkeypatch.setattr(shutil, "which", lambda name: "C:\\Program Files\\LiteLLM\\lite.EXE") - - helper = resolve_api_key_helper("https://gateway.example.com/?a=1&b=2", platform="win32") - - assert helper == ( - '"C:\\Program Files\\LiteLLM\\lite.EXE" "--base-url" "https://gateway.example.com/?a=1&b=2" ' - '"auth" "print-token"' - ) - - def test_non_windows_platforms_keep_posix_quoting(self, monkeypatch): - monkeypatch.setattr(shutil, "which", lambda name: "/usr/local/bin/lite") - - helper = resolve_api_key_helper("http://example.com/path; rm -rf /", platform="darwin") - - assert helper == "/usr/local/bin/lite --base-url 'http://example.com/path; rm -rf /' auth print-token" - - def _make_ctx(base_url): return click.Context(click.Command("test"), obj={"base_url": base_url}) @@ -307,7 +284,8 @@ def _capture_login(monkeypatch, on_login=lambda: None): login_calls = [] @click.pass_context - def fake_login(ctx, pkce=False): + def fake_login(ctx, config_claude=False, pkce=False): + assert config_claude is False, "`lite up` patches settings itself; the login it starts must not also configure" login_calls.append((ctx.obj["base_url"], pkce)) on_login() @@ -317,8 +295,8 @@ def _capture_login(monkeypatch, on_login=lambda: None): class TestEnsureFreshLogin: """A token that is fresh but was issued for a *different* proxy must not be trusted: without - this check, a user logged into proxy A who runs `up --base-url proxy-b` would silently get an - apiKeyHelper wired up around proxy A's real token, which print-token would then hand to proxy B.""" + this check, a user logged into proxy A who runs `up --base-url proxy-b` would silently get proxy A's + real token written into settings pointed at proxy B.""" def test_reuses_a_fresh_token_issued_for_the_same_proxy(self, monkeypatch): _FakeTokenStore( @@ -500,11 +478,13 @@ class TestUpCommand: settings_path, backup_path = _patch_paths(monkeypatch, tmp_path) original = {"theme": "dark"} settings_path.write_text(json.dumps(original)) + settings_path.chmod(0o644) captured = {} def fake_wait(self, timeout=None): captured["settings"] = json.loads(settings_path.read_text()) + captured["settings_mode"] = stat.S_IMODE(settings_path.stat().st_mode) captured["backup_existed"] = backup_path.exists() return True @@ -514,10 +494,6 @@ class TestUpCommand: patch(f"{UP_MODULE}.is_cli_token_fresh", return_value=True), patch(f"{UP_MODULE}.resolve_api_key", return_value="sk-fresh"), patch(f"{UP_MODULE}.verify_proxy_key"), - patch( - f"{UP_MODULE}.resolve_api_key_helper", - return_value="/usr/local/bin/lite auth print-token", - ), patch(f"{UP_MODULE}.signal.signal"), patch(f"{UP_MODULE}.atexit.register"), patch("threading.Event.wait", new=fake_wait), @@ -529,7 +505,11 @@ class TestUpCommand: assert captured["settings"]["theme"] == "dark" assert captured["settings"]["env"]["ANTHROPIC_BASE_URL"] == "http://localhost:4000" assert captured["settings"]["env"]["ENABLE_TOOL_SEARCH"] == "true" - assert captured["settings"]["apiKeyHelper"] == "/usr/local/bin/lite auth print-token" + assert captured["settings"]["env"]["ANTHROPIC_AUTH_TOKEN"] == "sk-fresh" + # The file now carries the key, so the umask (and the file's earlier 0644) must not decide who reads it. + assert captured["settings_mode"] == 0o600 + assert "apiKeyHelper" not in captured["settings"] + assert captured["settings"]["statusLine"]["command"].endswith("statusline.py") assert json.loads(settings_path.read_text()) == original assert not backup_path.exists() @@ -598,7 +578,7 @@ class TestUpCanInvokeTheRealLoginCommand: @click.pass_context def driver(ctx): ctx.obj = {"base_url": "http://127.0.0.1:9"} - ctx.invoke(real_login, pkce=False) + ctx.invoke(real_login, config_claude=False, pkce=False) with patch( f"{AUTH_MODULE}._start_cli_sso_flow", @@ -618,7 +598,7 @@ class TestUpCanInvokeTheRealLoginCommand: @click.pass_context def driver(ctx): ctx.obj = {"base_url": "http://127.0.0.1:9"} - ctx.invoke(real_login, pkce=False) + ctx.invoke(real_login, config_claude=False, pkce=False) with patch(f"{AUTH_MODULE}._start_cli_sso_flow", side_effect=RuntimeError("stop")): CliRunner().invoke(driver, [], standalone_mode=False, env={"CLAUDE_CONFIG_DIR": str(tmp_path)}) diff --git a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py index 7d5fc1a3544..4e2059ac30b 100644 --- a/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py +++ b/tests/test_litellm/proxy/common_utils/test_auth_cache_invalidation_pubsub.py @@ -1,4 +1,5 @@ import asyncio +import hashlib import json from typing import Iterable, List, Optional, Tuple from unittest.mock import patch @@ -7,6 +8,7 @@ import pytest from redis.asyncio import Redis from litellm.caching.in_memory_cache import InMemoryCache +from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.auth_cache_invalidation_pubsub import ( AUTH_CACHE_INVALIDATION_CHANNEL, AuthCacheInvalidationSubscriber, @@ -145,6 +147,35 @@ async def test_subscriber_deletes_local_cache_entry_on_message() -> None: assert pubsub.subscribed_channels == [AUTH_CACHE_INVALIDATION_CHANNEL] +@pytest.mark.asyncio +async def test_subscriber_deletes_key_object_partition_entry_on_message() -> None: + """ + LIT-7563 moved user-key objects into their own in-memory partition; a key + invalidation broadcast must still evict the hashed-token entry there, or a + deleted key keeps authenticating on other workers until its TTL expires. + """ + hashed_token = hashlib.sha256(b"sk-lit7563-hot-key").hexdigest() + cache = UserApiKeyCache() + cache.set_cache(hashed_token, UserAPIKeyAuth(token=hashed_token), model_type=UserAPIKeyAuth) + assert cache.get_cache(hashed_token, model_type=UserAPIKeyAuth) is not None + + pubsub = _QueuePubSub(initial_messages=[_invalidation_message(hashed_token)]) + subscriber = AuthCacheInvalidationSubscriber( + redis_cache=_FakeRedisCache(client=_ScriptedPubSubRedisClient(pubsubs=[pubsub])), + user_api_key_cache=cache, + ) + subscriber.start() + try: + for _ in range(200): + if cache.get_cache(hashed_token, model_type=UserAPIKeyAuth) is None: + break + await asyncio.sleep(0.01) + finally: + await subscriber.stop() + + assert cache.get_cache(hashed_token, model_type=UserAPIKeyAuth) is None + + @pytest.mark.asyncio async def test_subscriber_deletes_additional_in_memory_cache_entry_on_message() -> None: """ diff --git a/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py b/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py index 33e0d8bf38e..2d5d76ed542 100644 --- a/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py +++ b/tests/test_litellm/proxy/common_utils/test_user_api_key_cache.py @@ -1,3 +1,4 @@ +import hashlib import json from typing import Any @@ -10,10 +11,14 @@ from litellm.constants import DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL from litellm.proxy._types import UserAPIKeyAuth from litellm.proxy.common_utils.user_api_key_cache import ( UserApiKeyCache, + end_user_cache_key, get_management_object_ttl, + is_user_key_cache_key, ) from litellm.proxy.proxy_server import UserAPIKeyCacheTTLEnum +HASHED_TOKEN = hashlib.sha256(b"sk-lit7563-hot-key").hexdigest() + class CapturingInMemoryCache(InMemoryCache): """Records ``ttl`` passed into ``set_cache`` (what DualCache injects).""" @@ -204,9 +209,7 @@ class TestUserApiKeyCache: # Bypass UserApiKeyCache.serialize: CacheCodec rejects non-dict cached values # for dict-based models (deserialize returns None). - await cache.in_memory_cache.async_set_cache( - key="k", value="invalid-payload-not-a-dict" - ) + await cache.in_memory_cache.async_set_cache(key="k", value="invalid-payload-not-a-dict") value = await cache.async_get_cache("k", model_type=UserAPIKeyAuth) assert value is None @@ -224,6 +227,141 @@ class TestUserApiKeyCache: fake.set_cache("k2", {"ok": NotSerializable()}) +class TestUserKeyObjectPartition: + """ + Regression for LIT-7563: user-key objects share one 200-entry ``InMemoryCache`` with + every other management object, so end-user / team / tag churn evicts hot keys and + forces a ``LiteLLM_VerificationToken`` lookup on the next request. + """ + + @pytest.mark.parametrize( + ("key", "expected"), + [ + (HASHED_TOKEN, True), + (HASHED_TOKEN.upper(), False), + (f"team_id:{HASHED_TOKEN}", False), + (end_user_cache_key("u1"), False), + ("sk-lit7563-hot-key", False), + ], + ) + def test_is_user_key_cache_key(self, key: str, expected: bool): + assert is_user_key_cache_key(key) is expected + + @pytest.mark.asyncio + async def test_management_object_churn_does_not_evict_key_object(self): + cache = UserApiKeyCache(in_memory_cache=InMemoryCache(max_size_in_memory=2)) + await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth, ttl=100) + for i in range(2): + await cache.async_set_cache(end_user_cache_key(f"u{i}"), {"user_id": f"u{i}"}, ttl=200) + + key_obj = await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) + assert key_obj is not None + assert key_obj.token == HASHED_TOKEN + assert cache.get_cache(end_user_cache_key("u1")) == {"user_id": "u1"} + assert HASHED_TOKEN not in cache.in_memory_cache.cache_dict + + def test_sync_write_and_read_route_to_key_object_partition(self): + cache = UserApiKeyCache(in_memory_cache=InMemoryCache(max_size_in_memory=2)) + cache.set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth, ttl=100) + for i in range(2): + cache.set_cache(end_user_cache_key(f"u{i}"), {"user_id": f"u{i}"}, ttl=200) + + key_obj = cache.get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) + assert key_obj is not None + assert key_obj.token == HASHED_TOKEN + + @pytest.mark.asyncio + async def test_redis_hit_backfills_key_object_partition_with_configured_ttl(self): + redis = FakeRedisCache() + writer = UserApiKeyCache(redis_cache=redis, default_in_memory_ttl=30) + await writer.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth) + + key_partition = CapturingInMemoryCache() + reader = UserApiKeyCache(redis_cache=redis, default_in_memory_ttl=30, key_object_in_memory_cache=key_partition) + key_obj = await reader.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) + + assert key_obj is not None + assert key_obj.token == HASHED_TOKEN + assert key_partition.last_ttl == 30 + assert HASHED_TOKEN not in reader.in_memory_cache.cache_dict + + @pytest.mark.asyncio + async def test_update_cache_ttl_applies_to_key_object_partition(self): + key_partition = CapturingInMemoryCache() + cache = UserApiKeyCache(default_in_memory_ttl=60, key_object_in_memory_cache=key_partition) + cache.update_cache_ttl(default_in_memory_ttl=7, default_redis_ttl=7) + + await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth) + + assert key_partition.last_ttl == 7 + + @pytest.mark.asyncio + async def test_attach_redis_cache_applies_to_key_object_partition(self): + redis = FakeRedisCache() + cache = UserApiKeyCache() + cache.attach_redis_cache(redis) + + await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth) + + other_worker = UserApiKeyCache(redis_cache=redis) + key_obj = await other_worker.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) + assert key_obj is not None + assert key_obj.token == HASHED_TOKEN + + @pytest.mark.asyncio + async def test_delete_removes_key_object_from_partition_and_redis(self): + redis = FakeRedisCache() + cache = UserApiKeyCache(redis_cache=redis) + await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth) + assert await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) is not None + + cache.delete_cache(HASHED_TOKEN) + + assert await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) is None + assert await redis.async_get_cache(HASHED_TOKEN) is None + + @pytest.mark.asyncio + async def test_async_delete_removes_key_object_from_partition_and_redis(self): + redis = FakeRedisCache() + cache = UserApiKeyCache(redis_cache=redis) + await cache.async_set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth) + + await cache.async_delete_cache(HASHED_TOKEN) + + assert await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) is None + assert await redis.async_get_cache(HASHED_TOKEN) is None + + @pytest.mark.asyncio + async def test_pipeline_write_routes_each_entry_to_its_partition(self): + cache = UserApiKeyCache(in_memory_cache=InMemoryCache(max_size_in_memory=2)) + await cache.async_set_cache_pipeline( + [(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN))] + + [(end_user_cache_key(f"u{i}"), {"user_id": f"u{i}"}) for i in range(2)], + ttl=100, + ) + + key_obj = await cache.async_get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) + assert key_obj is not None + assert key_obj.token == HASHED_TOKEN + assert HASHED_TOKEN not in cache.in_memory_cache.cache_dict + assert cache.get_cache(end_user_cache_key("u1")) == {"user_id": "u1"} + + def test_flush_clears_key_object_partition(self): + cache = UserApiKeyCache() + cache.set_cache(HASHED_TOKEN, _make_key_obj(HASHED_TOKEN), model_type=UserAPIKeyAuth) + cache.set_cache(end_user_cache_key("u1"), {"user_id": "u1"}) + + cache.flush_cache() + + assert cache.get_cache(HASHED_TOKEN, model_type=UserAPIKeyAuth) is None + assert cache.get_cache(end_user_cache_key("u1")) is None + + def test_in_memory_cache_for_routes_by_key(self): + cache = UserApiKeyCache() + assert cache.in_memory_cache_for(HASHED_TOKEN) is cache.key_object_cache.in_memory_cache + assert cache.in_memory_cache_for(end_user_cache_key("u1")) is cache.in_memory_cache + + class TestManagementObjectTTL: """ Regression for LIT-3338: ``general_settings.user_api_key_cache_ttl`` (which the @@ -238,19 +376,13 @@ class TestManagementObjectTTL: def test_falls_back_to_constant_when_no_default_configured(self): cache = UserApiKeyCache() assert cache.default_in_memory_ttl is None - assert ( - get_management_object_ttl(cache) - == DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL - ) + assert get_management_object_ttl(cache) == DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL def test_resolves_on_a_plain_dual_cache(self): # Many call sites are typed UserApiKeyCache but exercised in tests with a # bare DualCache; the resolver must work on the base type, not just the subclass. assert get_management_object_ttl(DualCache(default_in_memory_ttl=300)) == 300 - assert ( - get_management_object_ttl(DualCache()) - == DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL - ) + assert get_management_object_ttl(DualCache()) == DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL @pytest.mark.asyncio async def test_management_write_uses_configured_ttl_over_constant(self): @@ -260,9 +392,7 @@ class TestManagementObjectTTL: redis_cache=FakeRedisCache(), default_in_memory_ttl=300, ) - assert get_management_object_ttl(cache) != ( - DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL - ) + assert get_management_object_ttl(cache) != (DEFAULT_MANAGEMENT_OBJECT_IN_MEMORY_CACHE_TTL) await cache.async_set_cache( "team_id:abc", diff --git a/tests/test_litellm/proxy/db/test_autorouter_session_rollup.py b/tests/test_litellm/proxy/db/test_autorouter_session_rollup.py index 4507892bd0f..271751a3ff8 100644 --- a/tests/test_litellm/proxy/db/test_autorouter_session_rollup.py +++ b/tests/test_litellm/proxy/db/test_autorouter_session_rollup.py @@ -59,7 +59,8 @@ class TestBuildTransaction: def test_successful_auto_routed_turn_builds_every_field(self): transaction = _build( metadata=_metadata( - usage_object={"prompt_tokens": 90, "cache_read_input_tokens": 5, "cache_creation_input_tokens": 7} + routing_decision={**ROUTING_DECISION, "savings_baseline_model": "anthropic/claude-opus-5"}, + usage_object={"prompt_tokens": 90, "cache_read_input_tokens": 5, "cache_creation_input_tokens": 7}, ) ) assert transaction == AutoRouterTurnTransaction( @@ -77,6 +78,7 @@ class TestBuildTransaction: cache_hit=True, cache_ttl_seconds=300, cache_touched=True, + baseline_model="anthropic/claude-opus-5", ) @pytest.mark.parametrize( @@ -109,6 +111,17 @@ class TestBuildTransaction: transaction = _build() assert transaction is not None and transaction.tier is None + def test_the_baseline_the_turn_was_priced_against_travels_with_the_turn(self): + decision = {**ROUTING_DECISION, "savings_baseline_model": "anthropic/claude-opus-5"} + transaction = _build(metadata=_metadata(routing_decision=decision)) + assert transaction is not None and transaction.baseline_model == "anthropic/claude-opus-5" + + @pytest.mark.parametrize("baseline", [None, "", 3]) + def test_a_decision_without_a_usable_baseline_records_none(self, baseline: object): + decision = {**ROUTING_DECISION, "savings_baseline_model": baseline} + transaction = _build(metadata=_metadata(routing_decision=decision)) + assert transaction is not None and transaction.baseline_model is None + def test_a_priced_classifier_rides_the_turns_spend(self): """The classifier row is excluded from the rollup, so its charge lands here, folded once into the turn that paid for it (GH #38816).""" @@ -215,6 +228,7 @@ def _transaction( session_id: str = "s1", at: datetime = datetime(2026, 8, 1, 12, 0, 0), tier: str | None = "medium", + baseline_model: str | None = "anthropic/claude-opus-5", ) -> AutoRouterTurnTransaction: return AutoRouterTurnTransaction( api_key="k1", @@ -232,6 +246,7 @@ def _transaction( cache_ttl_seconds=None, cache_touched=False, tier=tier, + baseline_model=baseline_model, ) @@ -265,6 +280,7 @@ class TestFlush: None, 0, "medium", + "anthropic/claude-opus-5", ) def test_a_connect_error_retries_the_same_statement(self): diff --git a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py index be55ac47bde..130b0da000b 100644 --- a/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py +++ b/tests/test_litellm/proxy/guardrails/guardrail_hooks/content_filter/test_content_filter.py @@ -626,6 +626,64 @@ class TestContentFilterGuardrail: assert entry["guardrail_status"] == "success" assert entry["guardrail_response"] == [] + @pytest.mark.asyncio + async def test_streaming_hook_duration_excludes_provider_wait(self): + """ + Streaming post-call: the logged guardrail duration must only cover the + per-chunk scans, not the time spent waiting on the provider between + chunks. PrometheusLogger adds post_call guardrail duration to + litellm_overhead_with_guardrails_latency_metric, so a duration spanning + the whole stream reports LLM generation time as guardrail overhead. + """ + import asyncio + + from litellm.types.utils import Delta, ModelResponseStream, StreamingChoices + + guardrail = ContentFilterGuardrail( + guardrail_name="test-streaming-duration", + patterns=[ + ContentFilterPattern( + pattern_type="prebuilt", + pattern_name="email", + action=ContentFilterAction.MASK, + ), + ], + event_hook=GuardrailEventHooks.post_call, + ) + + provider_wait_per_chunk = 0.15 + chunks = ("Hello ", "world, reach me at ", "test@example.com ") + + async def slow_stream(): + for i, text in enumerate(chunks): + await asyncio.sleep(provider_wait_per_chunk) + yield ModelResponseStream( + id=f"chunk{i}", + choices=[ + StreamingChoices( + delta=Delta(content=text), + index=0, + finish_reason="stop" if i == len(chunks) - 1 else None, + ) + ], + model="gpt-4", + ) + + request_data = {"messages": [{"role": "user", "content": "Hi"}], "model": "gpt-4o", "metadata": {}} + + async for _ in guardrail.async_post_call_streaming_iterator_hook( + user_api_key_dict=MagicMock(), + response=slow_stream(), + request_data=request_data, + ): + pass + + entry = request_data["metadata"]["standard_logging_guardrail_information"][0] + stream_wall_clock = entry["end_time"] - entry["start_time"] + assert stream_wall_clock >= provider_wait_per_chunk * len(chunks) + assert entry["masked_entity_count"].get("email", 0) >= 1 + assert 0 < entry["duration"] < provider_wait_per_chunk, entry["duration"] + @pytest.mark.asyncio async def test_streaming_hook_logs_guardrail_information_mask(self): """ diff --git a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py index dc18e0f7d4a..ef843adad98 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_auto_router_endpoints.py @@ -844,6 +844,143 @@ from litellm.proxy.management_endpoints.auto_router_endpoints import ( from litellm.types.management_endpoints.auto_router_endpoints import SHADOW_EVAL_TURN_VALVE, StartShadowEvalRequest VIEWER = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, api_key="sk-view", user_id="viewer") + + +class TestAutoRouterSession: + """GET /auto_router/session: a key reads its own session's routed model and savings, nothing else.""" + + ROW = { + "router_name": "claude-auto", + "router_type": "complexity", + "first_turn_at": datetime(2026, 9, 1, 12, 0, 0), + "last_turn_at": datetime(2026, 9, 1, 12, 5, 0), + "turns": 3, + "last_model": "anthropic/claude-sonnet-5", + "spend": 0.14, + "saved_spend": 0.24, + "classifier_cost": 0.0, + "tier_turns": {"simple": 1, "complex": 2}, + "baseline_models": {"anthropic/claude-opus-5": 3}, + } + + @staticmethod + def _rig(monkeypatch: pytest.MonkeyPatch, rows: Sequence[Mapping[str, object]]): + from litellm.proxy import proxy_server + + lookups: list[tuple[Mapping[str, object], Mapping[str, object]]] = [] + + class _Table: + async def find_first(self, where: Mapping[str, object], order: Mapping[str, object]): + lookups.append((where, order)) + matching = [r for r in rows if (r["api_key"], r["session_id"]) == (where["api_key"], where["session_id"])] + return max(matching, key=lambda r: r["last_turn_at"], default=None) + + monkeypatch.setattr( + proxy_server, "prisma_client", type("P", (), {"db": type("D", (), {"litellm_autoroutersession": _Table()})()})() + ) + return lookups + + @pytest.mark.asyncio + async def test_a_key_reads_its_own_session_with_the_baseline_its_turns_were_priced_against( + self, monkeypatch: pytest.MonkeyPatch + ): + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session + + caller = UserAPIKeyAuth(api_key="sk-caller") + self._rig(monkeypatch, [{**self.ROW, "api_key": caller.api_key, "session_id": "sess-1"}]) + response = await get_auto_router_session(user_api_key_dict=caller, session_id="sess-1") + assert response.model_dump() == { + "session_id": "sess-1", + "router_name": "claude-auto", + "router_type": "complexity", + "turns": 3, + "last_model": "anthropic/claude-sonnet-5", + "spend": 0.14, + "saved_spend": 0.24, + "baseline_spend": pytest.approx(0.38), + "baseline_model": "anthropic/claude-opus-5", + "baseline_models": {"anthropic/claude-opus-5": 3}, + } + + @pytest.mark.asyncio + async def test_another_keys_session_is_a_404_even_for_an_admin(self, monkeypatch: pytest.MonkeyPatch): + # The scope is the caller's own key hash, exactly what the spend writer keyed the row under; + # an admin wanting every key's sessions has /auto_router/benchmarks. + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session + + other = UserAPIKeyAuth(api_key="sk-other") + lookups = self._rig(monkeypatch, [{**self.ROW, "api_key": other.api_key, "session_id": "sess-1"}]) + with pytest.raises(HTTPException) as err: + await get_auto_router_session(user_api_key_dict=ADMIN, session_id="sess-1") + assert err.value.status_code == 404 + assert lookups == [({"api_key": ADMIN.api_key, "session_id": "sess-1"}, {"last_turn_at": "desc"})] + assert ADMIN.api_key != "sk-test" + + @pytest.mark.asyncio + async def test_the_sessions_most_recently_active_router_is_the_one_reported(self, monkeypatch: pytest.MonkeyPatch): + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session + + older = {**self.ROW, "api_key": ADMIN.api_key, "session_id": "s", "router_name": "old-auto"} + newer = { + **self.ROW, + "api_key": ADMIN.api_key, + "session_id": "s", + "router_name": "new-auto", + "last_turn_at": datetime(2026, 9, 1, 13, 0, 0), + } + self._rig(monkeypatch, [older, newer]) + response = await get_auto_router_session(user_api_key_dict=ADMIN, session_id="s") + assert response.router_name == "new-auto" + + @pytest.mark.asyncio + async def test_a_reconfigured_router_keeps_the_label_the_money_was_priced_against( + self, monkeypatch: pytest.MonkeyPatch + ): + # The proxy's router now prices against a different baseline, but the row's money was priced + # against opus for two of three turns, and the label says so; the full split is on the response. + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session + + priced = {"anthropic/claude-opus-5": 2, "anthropic/claude-sonnet-5": 1} + self._rig(monkeypatch, [{**self.ROW, "api_key": ADMIN.api_key, "session_id": "s", "baseline_models": priced}]) + response = await get_auto_router_session(user_api_key_dict=ADMIN, session_id="s") + assert response.baseline_model == "anthropic/claude-opus-5" + assert response.baseline_models == priced + + @pytest.mark.asyncio + async def test_a_session_whose_turns_recorded_no_baseline_reports_the_money_without_a_name( + self, monkeypatch: pytest.MonkeyPatch + ): + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session + + self._rig(monkeypatch, [{**self.ROW, "api_key": ADMIN.api_key, "session_id": "s", "baseline_models": {}}]) + response = await get_auto_router_session(user_api_key_dict=ADMIN, session_id="s") + assert response.baseline_model is None + assert response.baseline_spend == pytest.approx(0.38) + + @pytest.mark.asyncio + async def test_an_oversized_client_session_id_is_bounded_like_the_writer_bounded_it( + self, monkeypatch: pytest.MonkeyPatch + ): + from litellm.proxy.db.autorouter_session_rollup import bounded_session_id + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session + + long_id = "s" * 300 + self._rig(monkeypatch, [{**self.ROW, "api_key": ADMIN.api_key, "session_id": bounded_session_id(long_id)}]) + response = await get_auto_router_session(user_api_key_dict=ADMIN, session_id=long_id) + assert response.session_id == long_id + assert response.turns == 3 + + @pytest.mark.asyncio + async def test_without_a_database_the_endpoint_says_so(self, monkeypatch: pytest.MonkeyPatch): + from litellm.proxy import proxy_server + from litellm.proxy.management_endpoints.auto_router_endpoints import get_auto_router_session + + monkeypatch.setattr(proxy_server, "prisma_client", None) + with pytest.raises(HTTPException) as err: + await get_auto_router_session(user_api_key_dict=ADMIN, session_id="s") + assert err.value.status_code == 500 + + NON_ADMIN = UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, api_key="sk-user", user_id="user") diff --git a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py index addc952af14..7b285674145 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/test_llm_pass_through_endpoints.py @@ -4,7 +4,7 @@ import contextlib import json import os import traceback -from collections.abc import Mapping +from collections.abc import Iterator, Mapping from types import MappingProxyType, SimpleNamespace from typing import Final from unittest import mock @@ -12,7 +12,9 @@ from unittest.mock import AsyncMock, MagicMock, Mock, patch import httpx import pytest +import respx from fastapi import HTTPException, Request, Response +from fastapi.routing import APIRoute from fastapi.responses import StreamingResponse from fastapi.testclient import TestClient from starlette.datastructures import FormData @@ -45,6 +47,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import ( ) from litellm.proxy._types import LitellmUserRoles, SpecialHeaders, UserAPIKeyAuth from litellm.proxy.auth.handle_jwt import JWTHandler +from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials @@ -3339,8 +3342,11 @@ class TestOpenAIPassthroughRoute: def _resolve_route_name(method: str, path: str) -> str | None: from starlette.routing import Match + from litellm.proxy._lazy_features import LAZY_FEATURES, _force_load from litellm.proxy.proxy_server import app + asyncio.run(_force_load(app, next(f for f in LAZY_FEATURES if f.name == "llm_passthrough"))) + scope: Final = { "type": "http", "method": method, @@ -3350,8 +3356,8 @@ def _resolve_route_name(method: str, path: str) -> str | None: "root_path": "", } for route in app.router.routes: - if route.matches(scope)[0] == Match.FULL: - return getattr(route, "name", None) + if isinstance(route, APIRoute) and route.matches(scope)[0] == Match.FULL: + return route.name return None @@ -3376,7 +3382,7 @@ def test_openai_passthrough_prefix_wins_over_native_provider_routes(method, path /{provider}/v1/files and /{provider}/v1/batches routes must never capture it with provider="openai_passthrough" (which 500s on the LlmProviders lookup). """ - assert _resolve_route_name(method, path) == "openai_proxy_route" + assert _resolve_route_name(method, path) == "openai_passthrough_route" @pytest.mark.parametrize( @@ -3393,6 +3399,41 @@ def test_native_provider_routes_are_unchanged(method, path, expected_name): assert _resolve_route_name(method, path) == expected_name +@pytest.fixture +def openai_passthrough_client(monkeypatch: pytest.MonkeyPatch) -> Iterator[TestClient]: + from litellm.proxy.proxy_server import app + + monkeypatch.setenv("OPENAI_API_KEY", "sk-upstream") + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + monkeypatch.setitem(app.dependency_overrides, user_api_key_auth, lambda: UserAPIKeyAuth(api_key="sk-virtual")) + yield TestClient(app) + + +@pytest.mark.parametrize( + "method, path, body", + [ + ("POST", "/v1/responses", {"model": "gpt-5.1", "input": "hi"}), + ("GET", "/v1/files", None), + ("POST", "/v1/batches", {"input_file_id": "file-abc123", "endpoint": "/v1/responses"}), + ], +) +def test_openai_passthrough_forwards_verbatim_to_openai( + openai_passthrough_client: TestClient, method: str, path: str, body: dict[str, str] | None +) -> None: + """Every /openai_passthrough request, including the /v1/files and /v1/batches + paths that native provider routes also claim, must reach OpenAI unchanged.""" + with respx.mock(assert_all_called=True) as upstream: + route = upstream.request(method, f"https://api.openai.com{path}").mock( + return_value=httpx.Response(200, json={"id": "upstream_123"}) + ) + response = openai_passthrough_client.request(method, f"/openai_passthrough{path}", json=body) + + assert (response.status_code, response.json()) == (200, {"id": "upstream_123"}) + assert route.calls.last.request.headers["authorization"] == "Bearer sk-upstream" + + class TestCursorProxyRoute: """Tests for the Cursor Cloud Agents pass-through route.""" diff --git a/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py b/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py index cc231e383a3..87b1bc56659 100644 --- a/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py +++ b/tests/test_litellm/proxy/policy_engine/test_attachment_registry.py @@ -139,6 +139,71 @@ class TestGetAttachedPolicies: assert "gpt4-policy" in attached assert len(attached) == 3 + def test_matches_are_ordered_from_broadest_to_narrowest_scope(self): + registry = AttachmentRegistry() + registry.load_attachments( + [ + {"policy": "model-policy", "models": ["gpt-4"]}, + {"policy": "team-policy", "teams": ["t1"]}, + {"policy": "global-policy", "scope": "*"}, + ] + ) + + context = PolicyMatchContext(team_alias="t1", model="gpt-4") + + assert registry.get_attached_policies(context) == [ + "global-policy", + "team-policy", + "model-policy", + ] + + def test_combined_team_and_model_attachment_uses_model_specificity(self): + registry = AttachmentRegistry() + registry.load_attachments( + [ + {"policy": "team-policy", "teams": ["t1"]}, + {"policy": "team-model-policy", "teams": ["t1"], "models": ["gpt-4"]}, + ] + ) + + context = PolicyMatchContext(team_alias="t1", model="gpt-4") + + assert registry.get_attached_policies(context) == [ + "team-policy", + "team-model-policy", + ] + + def test_duplicate_policy_uses_broadest_matching_attachment(self): + registry = AttachmentRegistry() + registry.load_attachments( + [ + {"policy": "shared-policy", "models": ["gpt-4"]}, + {"policy": "model-policy", "models": ["gpt-4"]}, + {"policy": "shared-policy", "scope": "*"}, + ] + ) + + context = PolicyMatchContext(model="gpt-4") + + assert registry.get_attached_policies(context) == [ + "shared-policy", + "model-policy", + ] + assert registry.get_attached_policies_with_reasons(context)[0]["matched_via"] == "scope:*" + + def test_duplicate_policy_prefers_single_scope_over_combined_scope(self): + registry = AttachmentRegistry() + registry.load_attachments( + [ + {"policy": "shared-policy", "teams": ["t1"], "models": ["gpt-4"]}, + {"policy": "shared-policy", "models": ["gpt-4"]}, + ] + ) + + context = PolicyMatchContext(team_alias="t1", model="gpt-4") + + assert registry.get_attached_policies_with_reasons(context)[0]["matched_via"] == "model:gpt-4" + def test_same_policy_multiple_attachments_no_duplicates(self): """Test same policy attached multiple ways doesn't duplicate.""" registry = AttachmentRegistry() diff --git a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py index 300cac5cdb2..bc346874e0d 100644 --- a/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py +++ b/tests/test_litellm/proxy/proxy_server/test_team_model_name_translation.py @@ -27,6 +27,7 @@ from litellm.proxy.proxy_server import ( _get_proxy_model_info, _translate_model_name_for_response, ) +from litellm.types.router import DeploymentModelListingInfo def _team_row() -> dict: @@ -1089,6 +1090,101 @@ async def test_v1_models_metadata_does_not_leak_other_team_fallbacks(monkeypatch ] +@pytest.mark.asyncio +async def test_v1_models_team_alias_inherits_token_limits_and_chat_mode(monkeypatch): + team_dep = { + "model_name": "model_name_teamX_terra_uuid", + "litellm_params": {"model": "azure/gpt-4.1"}, + "model_info": { + "id": "id-terra", + "team_id": "teamX", + "team_public_model_name": "GPT Terra", + "access_groups": ["grp-a"], + "mode": "chat", + "max_input_tokens": 876000, + "max_output_tokens": 128000, + }, + } + router = MagicMock() + router.get_model_names.return_value = ["model_name_teamX_terra_uuid"] + router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_terra_uuid"]} + router.get_fully_blocked_model_names.return_value = set() + router.get_model_listing_info.return_value = DeploymentModelListingInfo( + cost_map_keys=("azure/gpt-4.1",), + max_input_tokens=876000, + max_output_tokens=128000, + ) + router.get_configured_mode.return_value = "chat" + router.model_list = [team_dep] + router.get_model_list.return_value = [team_dep] + router.get_model_group_info.return_value = None + + monkeypatch.setattr(ps, "llm_router", router) + monkeypatch.setattr(ps, "user_model", None) + monkeypatch.setattr(ps, "general_settings", {"use_team_public_model_name": True}) + + key = UserAPIKeyAuth(user_id="user", api_key="***", models=["grp-a"], team_models=[]) + response = await ps.model_list(user_api_key_dict=key, include_metadata=True) + + assert response["data"] == [ + { + "id": "GPT Terra", + "object": "model", + "created": 1677610602, + "owned_by": "openai", + "mode": "chat", + "max_input_tokens": 876000, + "max_output_tokens": 128000, + "metadata": {"fallbacks": []}, + } + ] + + +@pytest.mark.asyncio +async def test_v1_models_team_image_alias_inherits_image_generation_mode(monkeypatch): + team_dep = { + "model_name": "model_name_teamX_image_uuid", + "litellm_params": {"model": "openai/gpt-image-1"}, + "model_info": { + "id": "id-image", + "team_id": "teamX", + "team_public_model_name": "image", + "access_groups": ["grp-a"], + "mode": "image_generation", + }, + } + router = MagicMock() + router.get_model_names.return_value = ["model_name_teamX_image_uuid"] + router.get_model_access_groups.return_value = {"grp-a": ["model_name_teamX_image_uuid"]} + router.get_fully_blocked_model_names.return_value = set() + router.get_model_listing_info.return_value = DeploymentModelListingInfo( + cost_map_keys=("openai/gpt-image-1",), + max_input_tokens=None, + max_output_tokens=None, + ) + router.get_configured_mode.return_value = "image_generation" + router.model_list = [team_dep] + router.get_model_list.return_value = [team_dep] + router.get_model_group_info.return_value = None + + monkeypatch.setattr(ps, "llm_router", router) + monkeypatch.setattr(ps, "user_model", None) + monkeypatch.setattr(ps, "general_settings", {"use_team_public_model_name": True}) + + key = UserAPIKeyAuth(user_id="user", api_key="***", models=["grp-a"], team_models=[]) + response = await ps.model_list(user_api_key_dict=key) + + assert response["data"] == [ + { + "id": "image", + "object": "model", + "created": 1677610602, + "owned_by": "openai", + "mode": "image_generation", + } + ] + + def test_translate_team_model_names_for_listing_swaps_and_dedupes(): """Internal team routing keys -> public name; sibling deployments sharing a public name collapse to one entry (order preserved); globals untouched.""" @@ -1427,8 +1523,13 @@ async def test_retrieve_model_by_public_name_returns_200(monkeypatch): team_row = _team_row() router = _public_named_router(team_row) deployment = MagicMock() - deployment.litellm_params.model = "azure/gpt-5.2-low-rpm-testing" + deployment.litellm_params.model = "azure/gpt-4.1" router.get_deployment_by_model_group_name.return_value = deployment + router.get_model_listing_info.return_value = DeploymentModelListingInfo( + cost_map_keys=("azure/gpt-4.1",), + max_input_tokens=16384, + max_output_tokens=4096, + ) monkeypatch.setattr(ps, "llm_router", router) monkeypatch.setattr(ps, "general_settings", {}) @@ -1445,6 +1546,9 @@ async def test_retrieve_model_by_public_name_returns_200(monkeypatch): resp = await ps.model_info(model_id="team-claude-sonnet", user_api_key_dict=key) assert resp["id"] == "team-claude-sonnet" + assert resp.get("mode") == "chat" + assert resp.get("max_input_tokens") == 16384 + assert resp.get("max_output_tokens") == 4096 # lookup happened by the internal routing key, not the public name router.get_deployment_by_model_group_name.assert_called_once_with( "model_name_team-abc-123_4a6b8" diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py index 671a8ae63fc..6e43ac4a12b 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_management_endpoints.py @@ -6629,6 +6629,30 @@ def _session_page_row(session_key, last_activity): return {"session_key": session_key, "api_key": "hashed-key", "last_activity": last_activity} +def _session_grouped_paginating_prisma(sessions, counted_total=None): + """Mock prisma serving the grouped page query out of ``sessions``, honoring the LIMIT and OFFSET it asks for.""" + + async def mock_query_raw(sql_query, *params): + if "COUNT(*) AS total_count" in sql_query: + return [{"total_count": min(len(sessions) if counted_total is None else counted_total, params[-1])}] + if "DISTINCT ON" in sql_query: + return [_session_representative_row(f"req-{session_key}", session_key) for session_key in params[-2]] + if "COALESCE(SUM(spend)" in sql_query: + return [] + bounds = re.search(r"LIMIT \$(\d+)(?: OFFSET \$(\d+))?", sql_query) + limit = params[int(bounds.group(1)) - 1] + offset = params[int(bounds.group(2)) - 1] if bounds.group(2) else 0 + return [ + _session_page_row(session_key, last_activity) + for session_key, last_activity in sessions[offset : offset + limit] + ] + + mock_prisma = MagicMock() + mock_prisma.db = MagicMock() + mock_prisma.db.query_raw = AsyncMock(side_effect=mock_query_raw) + return mock_prisma + + @pytest.mark.asyncio async def test_ui_view_spend_logs_group_by_session_first_page(client, monkeypatch): """One row per (session, api_key), session-count total, and a keyset cursor for the next page.""" @@ -6742,6 +6766,203 @@ async def test_ui_view_spend_logs_group_by_session_cursor_page(client, monkeypat app.dependency_overrides.pop(ps.user_api_key_auth, None) +@pytest.mark.asyncio +async def test_ui_view_spend_logs_group_by_session_jumps_to_page_without_cursor(client, monkeypatch): + """page > 1 with no session_cursor (the UI's last-page jump) serves the sessions that page starts at.""" + sessions = tuple((f"sess-{index:02d}", f"2026-08-29 10:{59 - index:02d}:00") for index in range(60)) + mock_prisma = _session_grouped_paginating_prisma(sessions) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + try: + start_date, end_date = _default_date_range() + response = client.get( + "/spend/logs/ui", + params={ + "start_date": start_date, + "end_date": end_date, + "group_by_session": "true", + "page": 3, + "page_size": 25, + }, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 200, response.text + data = response.json() + assert data["page"] == 3 + assert data["total"] == 60 + assert data["has_more"] is False + assert data["next_session_cursor"] is None + assert [row["request_id"] for row in data["data"]] == [f"req-sess-{index:02d}" for index in range(50, 60)] + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_group_by_session_short_page_totals_itself(client, monkeypatch): + """A page that runs out of sessions is the end of the list, so the total comes from it and nothing is counted.""" + sessions = tuple((f"sess-{index:02d}", f"2026-08-29 10:{59 - index:02d}:00") for index in range(10)) + mock_prisma = _session_grouped_paginating_prisma(sessions, counted_total=999) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + try: + start_date, end_date = _default_date_range() + response = client.get( + "/spend/logs/ui", + params={ + "start_date": start_date, + "end_date": end_date, + "group_by_session": "true", + "page": 1, + "page_size": 25, + }, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 200, response.text + data = response.json() + assert data["total"] == 10, "the count query's 999 would have won if it had been asked" + assert data["total_is_capped"] is False + assert data["total_pages"] == 1 + assert len(data["data"]) == 10 + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_group_by_session_page_past_the_end_keeps_the_real_total(client, monkeypatch): + """An empty page past the last one says nothing about the total, so it is counted rather than inferred.""" + sessions = tuple((f"sess-{index:02d}", f"2026-08-29 10:{59 - index:02d}:00") for index in range(100)) + mock_prisma = _session_grouped_paginating_prisma(sessions) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + try: + start_date, end_date = _default_date_range() + response = client.get( + "/spend/logs/ui", + params={ + "start_date": start_date, + "end_date": end_date, + "group_by_session": "true", + "page": 4, + "page_size": 50, + }, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 200, response.text + data = response.json() + assert data["data"] == [] + assert data["total"] == 100, "the empty page's offset is not a total" + assert data["total_pages"] == 2 + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_group_by_session_page_past_count_cap_is_empty(client, monkeypatch): + """The last page inside the capped total still lists sessions; the page after it is empty and costs no query.""" + cap = spend_management_endpoints.SPEND_LOGS_PAGINATION_COUNT_CAP + sessions = tuple((f"sess-{index:06d}", "2026-08-29 10:00:00") for index in range(cap + 50)) + mock_prisma = _session_grouped_paginating_prisma(sessions) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + try: + start_date, end_date = _default_date_range() + params = { + "start_date": start_date, + "end_date": end_date, + "group_by_session": "true", + "page_size": 25, + } + last_page = client.get( + "/spend/logs/ui", + params={**params, "page": cap // 25}, + headers={"Authorization": "Bearer sk-test"}, + ) + assert last_page.status_code == 200, last_page.text + last_page_data = last_page.json() + assert last_page_data["total"] == cap + assert last_page_data["total_is_capped"] is True + assert last_page_data["data"][0]["request_id"] == f"req-sess-{cap - 25:06d}" + assert len(last_page_data["data"]) == 25 + + mock_prisma.db.query_raw.reset_mock() + past_cap = client.get( + "/spend/logs/ui", + params={**params, "page": cap // 25 + 1}, + headers={"Authorization": "Bearer sk-test"}, + ) + assert past_cap.status_code == 200, past_cap.text + past_cap_data = past_cap.json() + assert past_cap_data["data"] == [] + assert past_cap_data["has_more"] is False + assert past_cap_data["total"] == cap + assert mock_prisma.db.query_raw.await_count == 1, "only the bounded count query runs past the capped window" + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + +@pytest.mark.asyncio +async def test_ui_view_spend_logs_group_by_session_last_page_stops_at_the_capped_total(client, monkeypatch): + """A page size that does not divide the cap still ends the last page at the capped total it reports.""" + cap = spend_management_endpoints.SPEND_LOGS_PAGINATION_COUNT_CAP + sessions = tuple((f"sess-{index:06d}", "2026-08-29 10:00:00") for index in range(cap + 50)) + mock_prisma = _session_grouped_paginating_prisma(sessions) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) + monkeypatch.setattr( + "litellm.proxy.spend_tracking.spend_management_endpoints._is_admin_view_safe", + lambda user_api_key_dict: True, + ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin_user" + ) + try: + start_date, end_date = _default_date_range() + response = client.get( + "/spend/logs/ui", + params={ + "start_date": start_date, + "end_date": end_date, + "group_by_session": "true", + "page": cap // 7 + 1, + "page_size": 7, + }, + headers={"Authorization": "Bearer sk-test"}, + ) + assert response.status_code == 200, response.text + data = response.json() + assert data["total"] == cap + assert [row["request_id"] for row in data["data"]] == [ + f"req-sess-{index:06d}" for index in range(cap - cap % 7, cap) + ] + assert data["has_more"] is True + assert data["next_session_cursor"] is not None + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + @pytest.mark.asyncio async def test_ui_view_spend_logs_group_by_session_offset_for_non_starttime_sort( client, monkeypatch diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py index a7de3f1d8d6..54e5a6d5385 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_query_optimization.py @@ -528,8 +528,8 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch): group_key = "COALESCE(NULLIF(session_id, ''), request_id), api_key" session_rows = [ - {"session_key": "req-1", "api_key": "k", "last_activity": "2026-02-16 10:00:00"}, - {"session_key": "req-2", "api_key": "k", "last_activity": "2026-02-16 09:00:00"}, + {"session_key": f"req-{index}", "api_key": "k", "last_activity": f"2026-02-16 10:{59 - index:02d}:00"} + for index in range(51) ] representative_rows = [ {"request_id": "req-1", "api_key": "k", "metadata": "{}", "session_id": None}, @@ -538,7 +538,7 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch): async def mock_query_raw(sql_query, *params): if "COUNT(*) AS total_count" in sql_query: - return [{"total_count": 12}] + return [{"total_count": 60}] if "DISTINCT ON" in sql_query: return representative_rows return session_rows @@ -590,11 +590,11 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch): assert "COUNT(*) OVER ()" not in rep_sql assert [row["request_id"] for row in response["data"]] == ["req-1", "req-2"] - assert response["total"] == 12 + assert response["total"] == 60 assert response["total_is_capped"] is False - assert response["total_pages"] == 1 - assert response["has_more"] is False - assert response["next_session_cursor"] is None + assert response["total_pages"] == 2 + assert response["has_more"] is True + assert response["next_session_cursor"] == "2026-02-16 10:10:00|k|req-49" @pytest.mark.asyncio diff --git a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py index df113197ec6..dc8a87a97ad 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py +++ b/tests/test_litellm/proxy/spend_tracking/test_spend_tracking_utils.py @@ -1020,6 +1020,11 @@ _OVERLONG_MODEL: Final = "m" * (MAX_SPEND_LOG_MODEL_NAME_LENGTH + 1) ), ("gpt-5.2", ValueError("provider timed out"), "gpt-5.2"), (_BEDROCK_INFERENCE_PROFILE_ARN, ValueError("provider timed out"), _BEDROCK_INFERENCE_PROFILE_ARN), + ( + "MCP: deepwiki-ask_question", + ValueError("Content blocked: keyword 'confidential' detected"), + "MCP: deepwiki-ask_question", + ), ], ) def test_get_logging_payload_replaces_rejected_or_prompt_shaped_models_with_the_placeholder( diff --git a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py index 94aaa32519e..37b983d709a 100644 --- a/tests/test_litellm/proxy/test_litellm_pre_call_utils.py +++ b/tests/test_litellm/proxy/test_litellm_pre_call_utils.py @@ -23,6 +23,7 @@ from litellm.proxy.litellm_pre_call_utils import ( _get_dynamic_logging_metadata, _get_enforced_params, _get_metadata_variable_name, + _match_and_track_policies, _promoted_trace_control_fields, _resolve_credential_from_model_config, _resolve_provider_from_deployment, @@ -768,6 +769,7 @@ async def test_add_litellm_data_to_request_body_snapshot_excludes_proxy_server_r data = { "model": "gpt-3.5-turbo", "messages": [{"role": "user", "content": "hello"}], + "api_key": "request-key", } user_api_key_dict = UserAPIKeyAuth( @@ -795,6 +797,8 @@ async def test_add_litellm_data_to_request_body_snapshot_excludes_proxy_server_r assert "proxy_server_request" not in snapshot_body, ( "proxy_server_request must be excluded from its own body snapshot to prevent the body from self-referencing" ) + assert "api_key" not in snapshot_body + assert updated["proxy_server_request"]["credential_fields"] == ("api_key",) def test_refresh_proxy_server_request_body_snapshot_picks_up_guardrail_masking(): @@ -4149,6 +4153,30 @@ async def test_add_guardrails_from_policy_engine(): attachment_registry._initialized = False +def test_match_and_track_policies_preserves_attachment_and_request_body_order(): + from litellm.proxy.policy_engine.attachment_registry import AttachmentRegistry + from litellm.types.proxy.policy_engine import Policy, PolicyMatchContext + + attachment_policy_names = [f"attachment-policy-{index}" for index in range(8)] + request_body_policy_names = ["body-policy-1", "body-policy-2"] + policy_names = [*attachment_policy_names, *request_body_policy_names] + policies = {policy_name: Policy() for policy_name in policy_names} + attachment_registry = AttachmentRegistry() + attachment_registry.load_attachments( + [{"policy": policy_name, "scope": "*"} for policy_name in attachment_policy_names] + ) + + applied_policy_names, _ = _match_and_track_policies( + data={"metadata": {}}, + context=PolicyMatchContext(model="gpt-4"), + request_body_policies=request_body_policy_names, + policies_override=policies, + attachment_registry_override=attachment_registry, + ) + + assert applied_policy_names == policy_names + + @pytest.mark.asyncio async def test_add_guardrails_from_policy_engine_keeps_a_policy_added_guardrail_its_pipeline_also_steps(): from litellm.proxy.policy_engine.attachment_registry import get_attachment_registry diff --git a/tests/test_litellm/proxy/test_proxy_server.py b/tests/test_litellm/proxy/test_proxy_server.py index e058a4f6396..3b3031647dc 100644 --- a/tests/test_litellm/proxy/test_proxy_server.py +++ b/tests/test_litellm/proxy/test_proxy_server.py @@ -7313,6 +7313,88 @@ async def test_update_general_settings_apply_user_budget_to_team_keys_yaml_wins( assert ps.general_settings["apply_user_budget_to_team_keys"] is True +def _fill_user_api_key_cache(cache: DualCache, count: int) -> None: + for index in range(count): + cache.set_cache(key=f"key-{index}", value={"token": f"key-{index}"}, local_only=True) + + +@pytest.mark.asyncio +async def test_update_general_settings_user_api_key_cache_max_size_resizes_the_running_cache(monkeypatch): + """The Admin UI writes the capacity to the DB config, so the running cache has + to pick it up on reload; otherwise the knob only works after a restart.""" + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.proxy_server import ProxyConfig + + cache = UserApiKeyCache() + monkeypatch.setattr(proxy_server_module, "general_settings", {}) + monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache) + await ProxyConfig()._update_general_settings(db_general_settings={"user_api_key_cache_max_size": 300}) + + assert proxy_server_module.general_settings["user_api_key_cache_max_size"] == 300 + + _fill_user_api_key_cache(cache, 250) + assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"} + + +@pytest.mark.asyncio +async def test_update_general_settings_clearing_user_api_key_cache_max_size_restores_the_default(monkeypatch): + """Blanking the field in the dashboard deletes the key, so the cache must fall + back to the default capacity rather than keep the last configured size.""" + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.proxy_server import ProxyConfig + + cache = UserApiKeyCache() + cache.update_in_memory_max_size(5000) + monkeypatch.setattr(proxy_server_module, "general_settings", {"user_api_key_cache_max_size": 5000}) + monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache) + await ProxyConfig()._update_general_settings(db_general_settings={"store_model_in_db": True}) + + assert "user_api_key_cache_max_size" not in proxy_server_module.general_settings + + _fill_user_api_key_cache(cache, 201) + assert cache.get_cache(key="key-0", local_only=True) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("db_value", [0, -5, "lots"]) +async def test_update_general_settings_ignores_an_invalid_user_api_key_cache_max_size(db_value, monkeypatch): + """A non-positive capacity would make the eviction loop pop an empty heap on the + next write, so a bad DB value must leave the running cache untouched.""" + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.proxy_server import ProxyConfig + + cache = UserApiKeyCache() + cache.update_in_memory_max_size(300) + monkeypatch.setattr(proxy_server_module, "general_settings", {}) + monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache) + await ProxyConfig()._update_general_settings(db_general_settings={"user_api_key_cache_max_size": db_value}) + + assert "user_api_key_cache_max_size" not in proxy_server_module.general_settings + + _fill_user_api_key_cache(cache, 250) + assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"} + + +@pytest.mark.asyncio +async def test_update_general_settings_user_api_key_cache_max_size_yaml_wins(monkeypatch): + """A DB value must not silently override an explicit YAML capacity on reload.""" + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.proxy_server import ProxyConfig + + proxy_config = ProxyConfig() + proxy_config._yaml_general_settings_keys = {"user_api_key_cache_max_size"} + cache = UserApiKeyCache() + cache.update_in_memory_max_size(300) + monkeypatch.setattr(proxy_server_module, "general_settings", {"user_api_key_cache_max_size": 300}) + monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache) + await proxy_config._update_general_settings(db_general_settings={"user_api_key_cache_max_size": 10}) + + assert proxy_server_module.general_settings["user_api_key_cache_max_size"] == 300 + + _fill_user_api_key_cache(cache, 250) + assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"} + + @pytest.mark.asyncio @pytest.mark.parametrize( "db_value,expected", @@ -9183,6 +9265,126 @@ class TestLazyFeaturesNotImportedAtStartup: class TestLazyFeatureMiddleware: """Behavior of the middleware itself, exercised in isolation.""" + @pytest.mark.asyncio + async def test_llm_passthrough_loads_on_first_provider_request(self, monkeypatch): + """An app that never registered the provider passthrough routes 404s a + provider request; behind the middleware the same request registers the + routes and is forwarded to the provider with the configured key.""" + import respx + from fastapi import FastAPI + + from litellm.proxy._lazy_features import LAZY_FEATURES, LazyFeatureMiddleware + + monkeypatch.setenv("MISTRAL_API_KEY", "sk-upstream") + monkeypatch.delenv("SERVER_ROOT_PATH", raising=False) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + litellm.in_memory_llm_clients_cache.flush_cache() + feat = next(f for f in LAZY_FEATURES if f.name == "llm_passthrough") + target_app = FastAPI() + target_app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(api_key="sk-virtual") + mw = LazyFeatureMiddleware(target_app, fastapi_app=target_app, features=(feat,)) + + with respx.mock() as upstream: + route = upstream.get("https://api.mistral.ai/v1/models").mock( + return_value=httpx.Response(200, json={"object": "list", "data": []}) + ) + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=target_app), base_url="http://t") as bare: + assert (await bare.get("/mistral/v1/models")).status_code == 404 + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=mw), base_url="http://t") as lazy: + response = await lazy.get("/mistral/v1/models") + + assert (response.status_code, response.json()) == (200, {"object": "list", "data": []}) + assert route.calls.last.request.headers["authorization"] == "Bearer sk-upstream" + + def test_llm_passthrough_prefixes_cover_every_route_the_module_registers(self): + """A route the module registers under a prefix the feature does not claim + would 404 until an unrelated provider request happens to load the module.""" + from litellm.proxy._lazy_features import LAZY_FEATURES + + feat = next(f for f in LAZY_FEATURES if f.name == "llm_passthrough") + paths = [r.path for r in importlib.import_module(feat.module_path).router.routes] + + assert {"/mistral/{endpoint:path}", "/openai/{endpoint:path}"} <= set(paths) + unreachable = [p for p in paths if not feat.matches(p.replace("{endpoint:path}", "x"))] + assert unreachable == [], f"routes the middleware would never load: {unreachable}" + + @pytest.mark.asyncio + @pytest.mark.parametrize("first_hit", ["/v1/realtime/calls", "/openai/v1/models"]) + async def test_lazy_routes_land_in_registry_order_not_first_hit_order(self, first_hit): + """Two lazy features with overlapping paths must answer a request with the + same handler no matter which one a deployment happens to hit first.""" + from fastapi import APIRouter, FastAPI + + from litellm.proxy._lazy_features import LazyFeature, LazyFeatureMiddleware + + def make_register(path, handler): + def register(app, module): + router = APIRouter() + router.add_api_route(path, lambda: {"handler": handler}, methods=["POST"]) + app.include_router(router) + + return register + + catch_all = LazyFeature( + name="catch_all", + module_path="json", + path_prefixes=("/openai/",), + register_fn=make_register("/openai/{endpoint:path}", "catch_all"), + ) + specific = LazyFeature( + name="specific", + module_path="base64", + path_prefixes=("/openai/v1/realtime", "/v1/realtime"), + register_fn=make_register("/openai/v1/realtime/calls", "specific"), + ) + + target_app = FastAPI() + mw = LazyFeatureMiddleware(target_app, fastapi_app=target_app, features=(catch_all, specific)) + + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=mw), base_url="http://t") as client: + await client.post(first_hit) + await client.post("/openai/v1/models") + response = await client.post("/openai/v1/realtime/calls") + + assert response.json() == {"handler": "catch_all"} + + @pytest.mark.asyncio + @pytest.mark.parametrize("root_path", ["", "/api"]) + async def test_reserved_slot_keeps_lazy_catch_all_ahead_of_later_eager_routes(self, root_path): + """/{mcp_server_name}/mcp is registered after the provider passthrough router + at startup, so /mistral/mcp must keep reaching the provider catch-all once + that router loads lazily instead of being swallowed by the MCP route. The + native /mistral/v1/files route sits ahead of it, so that path neither loads + the feature nor changes owner, with or without a SERVER_ROOT_PATH prefix.""" + from fastapi import APIRouter, FastAPI + + from litellm.proxy._lazy_features import LazyFeature, LazyFeatureMiddleware, reserve_lazy_slot + + def register(app, module): + router = APIRouter() + router.add_api_route("/mistral/{endpoint:path}", lambda: {"handler": "passthrough"}, methods=["POST"]) + app.include_router(router) + + passthrough = LazyFeature( + name="llm_passthrough", module_path="json", path_prefixes=("/mistral/",), register_fn=register + ) + target_app = FastAPI(root_path=root_path) + target_app.add_api_route("/mistral/v1/files", lambda: {"handler": "files"}, methods=["POST"]) + reserve_lazy_slot(target_app, "llm_passthrough", features=(passthrough,)) + target_app.add_api_route("/{mcp_server_name}/mcp", lambda: {"handler": "mcp"}, methods=["POST"]) + target_app.add_middleware(LazyFeatureMiddleware, fastapi_app=target_app, features=(passthrough,)) + + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=target_app), base_url="http://t") as client: + files_first = (await client.post(f"{root_path}/mistral/v1/files")).json()["handler"] + loaded_after_files = frozenset(target_app.state.lazy_loaded) + handlers = [ + (await client.post(f"{root_path}{path}")).json()["handler"] + for path in ("/mistral/mcp", "/mistral/v1/files") + ] + + assert (files_first, loaded_after_files) == ("files", frozenset()) + assert handlers == ["passthrough", "files"] + @pytest.mark.asyncio async def test_first_request_triggers_load_subsequent_does_not(self): from fastapi import FastAPI @@ -10224,6 +10426,27 @@ def test_get_config_list_includes_apply_user_budget_to_team_keys(monkeypatch): app.dependency_overrides.clear() +def test_get_config_list_includes_user_api_key_cache_max_size(monkeypatch): + """The Admin UI General Settings table renders whatever /config/list returns, + so the cache capacity has to be exposed there as an Integer to be editable.""" + mock_prisma = MagicMock() + mock_config_table = MagicMock() + mock_config_table.find_first = AsyncMock(return_value=None) + mock_prisma.db = types.SimpleNamespace(litellm_config=mock_config_table) + monkeypatch.setattr(proxy_server_module, "prisma_client", mock_prisma) + app.dependency_overrides[proxy_server_module.user_api_key_auth] = lambda: UserAPIKeyAuth( + user_id="admin", user_role=LitellmUserRoles.PROXY_ADMIN + ) + try: + client = TestClient(app) + resp = client.get("/config/list", params={"config_type": "general_settings"}) + assert resp.status_code == 200, resp.text + fields = {item["field_name"]: item for item in resp.json()} + assert fields["user_api_key_cache_max_size"]["field_type"] == "Integer" + finally: + app.dependency_overrides.clear() + + def test_get_config_list_includes_budget_exceeded_throttle_percentage(monkeypatch): """The throttle fraction is a litellm_settings scalar surfaced on the General Settings table as a Float field so it sits with the other global limits; it @@ -12826,6 +13049,48 @@ async def test_load_config_router_authorizes_fallback_targets_against_the_callin assert router.fallback_access_check is router_fallback_access_check +@pytest.mark.asyncio +async def test_load_config_user_api_key_cache_max_size_keeps_more_than_200_entries(tmp_path, monkeypatch): + """The auth cache used to be pinned at InMemoryCache's 200 entry default, so a + deployment with more keys than that evicted constantly and every request + fell through to the DB. The YAML knob has to raise the cap on the live cache.""" + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.proxy_server import ProxyConfig + + config_file = tmp_path / "config.yaml" + config_file.write_text(yaml.dump({"general_settings": {"user_api_key_cache_max_size": "1000"}})) + + cache = UserApiKeyCache() + monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache) + await ProxyConfig().load_config(router=None, config_file_path=str(config_file)) + + _fill_user_api_key_cache(cache, 999) + assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"} + + +@pytest.mark.asyncio +@pytest.mark.parametrize("bad_value", [0, -1, "unbounded"]) +async def test_load_config_rejects_a_non_positive_user_api_key_cache_max_size(tmp_path, bad_value, monkeypatch): + """InMemoryCache treats 0 as 'cache nothing' and a negative cap makes eviction + pop an empty heap, so the proxy must refuse to boot with such a value instead + of silently disabling auth caching.""" + from pydantic import ValidationError + + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.proxy_server import ProxyConfig + + config_file = tmp_path / "config.yaml" + config_file.write_text(yaml.dump({"general_settings": {"user_api_key_cache_max_size": bad_value}})) + + cache = UserApiKeyCache() + monkeypatch.setattr(proxy_server_module, "user_api_key_cache", cache) + with pytest.raises(ValidationError): + await ProxyConfig().load_config(router=None, config_file_path=str(config_file)) + + _fill_user_api_key_cache(cache, 150) + assert cache.get_cache(key="key-0", local_only=True) == {"token": "key-0"} + + def test_docs_redoc_openapi_are_reachable_by_default(): """ LIT-6745: the interactive/machine-readable docs surfaces are on by diff --git a/tests/test_litellm/proxy/test_proxy_utils.py b/tests/test_litellm/proxy/test_proxy_utils.py index 646596c5b88..9def21c0573 100644 --- a/tests/test_litellm/proxy/test_proxy_utils.py +++ b/tests/test_litellm/proxy/test_proxy_utils.py @@ -581,6 +581,75 @@ class TestPostCallFailureHookLiftsStandardLoggingObject: assert "standard_logging_object" not in request_data +class TestPostCallFailureHookLiftsCallTypeAndStartTime: + """A guardrail-blocked MCP tool call fails before any LLM call. The failure + spend row is built from request_data after ``litellm_logging_obj`` is popped, + so ``call_type`` and the request ``start_time`` must be lifted off the logging + object first, or the Logs page shows the row as an LLM call with a blank call + type and a 0s duration (LIT-7453). + """ + + @pytest.mark.asyncio + async def test_failed_mcp_tool_call_spend_row_keeps_call_type_model_and_duration(self): + import traceback + from types import SimpleNamespace + from unittest.mock import AsyncMock + + from litellm.litellm_core_utils.litellm_logging import Logging + from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger + from litellm.proxy.spend_tracking.spend_tracking_utils import get_logging_payload + + request_start = real_datetime.datetime.now() - real_datetime.timedelta(seconds=2) + logging_obj = Logging( + model="MCP: deepwiki-ask_question", + messages=[], + stream=False, + call_type="call_mcp_tool", + start_time=request_start, + litellm_call_id="call-1", + function_id="fn-1", + ) + logging_obj.update_environment_variables( + model="MCP: deepwiki-ask_question", + user="", + optional_params={}, + litellm_params={"metadata": {"user_api_key_hash": "hashed"}}, + ) + blocked = Exception("Content blocked: keyword 'confidential' detected") + logging_obj.failure_handler(blocked, traceback.format_exc(), request_start, real_datetime.datetime.now()) + request_data = { + "name": "deepwiki-ask_question", + "arguments": {"question": "confidential"}, + "litellm_logging_obj": logging_obj, + } + proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache()) + proxy_logging_obj.alert_types = [] + spend_writer = SimpleNamespace(update_database=AsyncMock()) + original_callbacks = list(litellm.callbacks) + litellm.callbacks = [_ProxyDBLogger(spend_writer=lambda: spend_writer)] + try: + with patch.object(proxy_logging_obj, "update_request_status", new=AsyncMock()): + await proxy_logging_obj.post_call_failure_hook( + request_data=request_data, + original_exception=blocked, + user_api_key_dict=UserAPIKeyAuth(api_key="sk-test"), + ) + finally: + litellm.callbacks = original_callbacks + ProxyLogging._callback_capabilities_cache.clear() + + db_call = spend_writer.update_database.call_args.kwargs + payload = get_logging_payload( + kwargs=db_call["kwargs"], + response_obj=db_call["completion_response"], + start_time=db_call["start_time"], + end_time=db_call["end_time"], + ) + assert payload["call_type"] == "call_mcp_tool" + assert payload["model"] == "MCP: deepwiki-ask_question" + assert payload["endTime"] - payload["startTime"] >= real_datetime.timedelta(seconds=2) + + class TestPostCallFailureHookEstimatesDispatchedInputTokens: """A non-stream request that failed after dispatch (timeout, provider error) consumed provider-billed input tokens but recovered no usage. @@ -1848,13 +1917,14 @@ def test_a_failure_with_no_logging_object_lifts_nothing(): assert dict(_failure_fields_to_lift({"litellm_logging_obj": _LoggingObj({})})) == {} -def test_a_dispatched_failure_lifts_the_four_fields_the_spend_log_needs(): +def test_a_dispatched_failure_lifts_the_fields_the_spend_log_needs(): from litellm.proxy.utils import _failure_fields_to_lift lifted = _failure_fields_to_lift( { "litellm_logging_obj": _LoggingObj( { + "start_time": 1699999999.0, "first_api_call_start_time": 1700000000.0, "call_type": "acompletion", "model": FAILURE_USAGE_MODEL, @@ -1866,12 +1936,16 @@ def test_a_dispatched_failure_lifts_the_four_fields_the_spend_log_needs(): ) assert set(lifted) == { + "start_time", "first_api_call_start_time", + "call_type", "combined_usage_object", "response_cost", "standard_logging_object", } + assert lifted["start_time"] == 1699999999.0 assert lifted["first_api_call_start_time"] == 1700000000.0 + assert lifted["call_type"] == "acompletion" assert lifted["response_cost"] == 0.0 assert lifted["combined_usage_object"].prompt_tokens > 0 assert lifted["standard_logging_object"] == {"id": "log-1"} diff --git a/tests/test_litellm/proxy/test_route_priority.py b/tests/test_litellm/proxy/test_route_priority.py new file mode 100644 index 00000000000..dfdc816f4b4 --- /dev/null +++ b/tests/test_litellm/proxy/test_route_priority.py @@ -0,0 +1,171 @@ +import sys +from types import ModuleType + +import httpx +import pytest +from fastapi import APIRouter, FastAPI +from fastapi.testclient import TestClient +from starlette.routing import Match + +from litellm.proxy.route_priority import HOT_ROUTE_PATHS, hot_routes_first + +FILLER_COUNT = 300 + + +def _routes_scanned_before_dispatch(app: FastAPI, method: str, path: str) -> int: + """Number of route.matches() calls Starlette's Router.app makes before it finds a full match.""" + scope = {"type": "http", "method": method, "path": path, "root_path": "", "headers": [], "query_string": b""} + for i, route in enumerate(app.router.routes): + match, _ = route.matches(dict(scope)) + if match == Match.FULL: + return i + 1 + raise AssertionError(f"{method} {path} has no route") + + +def _hot_router() -> APIRouter: + router = APIRouter() + + @router.get("/health/liveliness") + @router.get("/health/liveness") + async def liveliness(): + return "I'm alive!" + + @router.post("/v1/chat/completions") + @router.post("/chat/completions") + async def chat(): + return {"object": "chat.completion"} + + return router + + +def _app_with_filler_then_hot_routes() -> FastAPI: + app = FastAPI() + for i in range(FILLER_COUNT): + + @app.get(f"/filler/{i}") + async def filler(i: int = i): + return {"filler": i} + + app.include_router(_hot_router()) + return app + + +def test_hot_routes_first_puts_hot_routes_ahead_of_everything_else(): + app = _app_with_filler_then_hot_routes() + assert _routes_scanned_before_dispatch(app, "GET", "/health/liveliness") > FILLER_COUNT + + app.router.routes = hot_routes_first(app.router.routes) + + hot_count = sum(1 for r in app.router.routes if getattr(r, "path", None) in HOT_ROUTE_PATHS) + assert _routes_scanned_before_dispatch(app, "GET", "/health/liveliness") <= hot_count + assert _routes_scanned_before_dispatch(app, "GET", "/health/liveness") <= hot_count + assert _routes_scanned_before_dispatch(app, "POST", "/v1/chat/completions") <= hot_count + assert _routes_scanned_before_dispatch(app, "POST", "/chat/completions") <= hot_count + + +def test_hot_routes_first_keeps_the_other_routes_in_order_and_dispatching(): + app = _app_with_filler_then_hot_routes() + before = [r.path for r in app.router.routes if getattr(r, "path", "").startswith("/filler/")] + + app.router.routes = hot_routes_first(app.router.routes) + + after = [r.path for r in app.router.routes if getattr(r, "path", "").startswith("/filler/")] + assert after == before + client = TestClient(app) + assert client.get("/health/liveliness").json() == "I'm alive!" + assert client.get("/filler/7").json() == {"filler": 7} + assert client.post("/v1/chat/completions").json() == {"object": "chat.completion"} + assert client.get("/v1/chat/completions").status_code == 405 + assert client.get("/does/not/exist").status_code == 404 + + +def test_hot_routes_first_is_idempotent(): + app = _app_with_filler_then_hot_routes() + once = hot_routes_first(app.router.routes) + assert hot_routes_first(once) == once + + +@pytest.mark.asyncio +async def test_lazy_loaded_hot_route_moves_to_the_front(monkeypatch): + from litellm.proxy._lazy_features import LazyFeature, LazyFeatureMiddleware + + messages_router = APIRouter() + + @messages_router.post("/v1/messages") + async def messages(): + return {"type": "message"} + + fake_module = ModuleType("fake_anthropic_endpoints") + fake_module.router = messages_router + monkeypatch.setitem(sys.modules, fake_module.__name__, fake_module) + + target_app = _app_with_filler_then_hot_routes() + target_app.router.routes = hot_routes_first(target_app.router.routes) + + async def downstream(scope, receive, send): + await send({"type": "http.response.start", "status": 200, "headers": []}) + await send({"type": "http.response.body", "body": b""}) + + feat = LazyFeature(name="anthropic", module_path=fake_module.__name__, path_prefixes=("/v1/messages",)) + mw = LazyFeatureMiddleware(downstream, fastapi_app=target_app, features=(feat,)) + + async def receive(): + return {"type": "http.request", "body": b"", "more_body": False} + + async def send(message): + pass + + await mw({"type": "http", "path": "/v1/messages", "method": "POST", "headers": []}, receive, send) + + hot_count = sum(1 for r in target_app.router.routes if getattr(r, "path", None) in HOT_ROUTE_PATHS) + assert _routes_scanned_before_dispatch(target_app, "POST", "/v1/messages") <= hot_count + assert TestClient(target_app).post("/v1/messages").json() == {"type": "message"} + + +@pytest.mark.asyncio +async def test_hot_routes_first_keeps_reserved_lazy_slot_ahead_of_later_eager_routes(): + """Liveness is registered after the provider passthrough slot, so pulling it to the + front must not shift where the lazily loaded catch-all is spliced back in.""" + from litellm.proxy._lazy_features import LazyFeature, LazyFeatureMiddleware, reserve_lazy_slot + + def register(app, module): + router = APIRouter() + router.add_api_route("/mistral/{endpoint:path}", lambda: {"handler": "passthrough"}, methods=["POST"]) + app.include_router(router) + + passthrough = LazyFeature( + name="llm_passthrough", module_path="json", path_prefixes=("/mistral/",), register_fn=register + ) + target_app = FastAPI() + target_app.add_api_route("/mistral/v1/files", lambda: {"handler": "files"}, methods=["POST"]) + target_app.add_api_route("/mistral/v1/batches", lambda: {"handler": "batches"}, methods=["POST"]) + reserve_lazy_slot(target_app, "llm_passthrough", features=(passthrough,)) + target_app.include_router(_hot_router()) + target_app.add_api_route("/{mcp_server_name}/mcp", lambda: {"handler": "mcp"}, methods=["POST"]) + target_app.router.routes = hot_routes_first(target_app.router.routes) + target_app.add_middleware(LazyFeatureMiddleware, fastapi_app=target_app, features=(passthrough,)) + + async with httpx.AsyncClient(transport=httpx.ASGITransport(app=target_app), base_url="http://t") as client: + batches_first = (await client.post("/mistral/v1/batches")).json()["handler"] + loaded_after_batches = frozenset(target_app.state.lazy_loaded) + handlers = [ + (await client.post(path)).json()["handler"] + for path in ("/mistral/mcp", "/mistral/v1/files", "/mistral/v1/batches") + ] + + assert (batches_first, loaded_after_batches) == ("batches", frozenset()) + assert handlers == ["passthrough", "files", "batches"] + hot_count = sum(1 for r in target_app.router.routes if getattr(r, "path", None) in HOT_ROUTE_PATHS) + assert _routes_scanned_before_dispatch(target_app, "GET", "/health/liveliness") <= hot_count + + +def test_proxy_app_dispatches_liveness_and_chat_completions_before_the_rest(): + from litellm.proxy.proxy_server import app + + hot_count = sum(1 for r in app.router.routes if getattr(r, "path", None) in HOT_ROUTE_PATHS) + assert hot_count >= 4 + assert len(app.router.routes) > 100 + assert _routes_scanned_before_dispatch(app, "GET", "/health/liveliness") <= hot_count + assert _routes_scanned_before_dispatch(app, "GET", "/health/liveness") <= hot_count + assert _routes_scanned_before_dispatch(app, "POST", "/v1/chat/completions") <= hot_count + assert _routes_scanned_before_dispatch(app, "POST", "/chat/completions") <= hot_count diff --git a/tests/test_litellm/repositories/test_repositories.py b/tests/test_litellm/repositories/test_repositories.py index a890d7ceed0..c7f6b8ff83a 100644 --- a/tests/test_litellm/repositories/test_repositories.py +++ b/tests/test_litellm/repositories/test_repositories.py @@ -2407,3 +2407,60 @@ class TestCountBillableUsers: client.db.litellm_usertable = _RacyTable() repo = UserRepository(client) assert await repo.count_billable_users() == 0 + + +class TestAutoRouterSessionRepository: + ROW: Final = { + "api_key": "hashed-key", + "session_id": "s1", + "router_name": "claude-auto", + "router_type": "complexity", + "first_turn_at": datetime(2026, 9, 1, 12, 0, 0), + "last_turn_at": datetime(2026, 9, 1, 12, 5, 0), + "last_model": "anthropic/claude-sonnet-5", + "models": {"anthropic/claude-sonnet-5": {"at": 1.0, "ttl": None}}, + "turns": 3, + "spend": 0.14, + "saved_spend": 0.24, + "classifier_cost": 0.01, + "tier_turns": {"complex": 3}, + "baseline_models": {"anthropic/claude-opus-5": 3}, + } + + @staticmethod + def _repo(record: Optional[Dict[str, Any]]): + from litellm.repositories.autorouter_session_repository import AutoRouterSessionRepository + + lookups: List[Dict[str, Any]] = [] + + class _Table: + async def find_first(self, where: Dict[str, Any], order: Dict[str, str]): + lookups.append({"where": where, "order": order}) + return MockRecord(record) if record is not None else None + + client = MagicMock() + client.db.litellm_autoroutersession = _Table() + return AutoRouterSessionRepository(client), lookups + + @pytest.mark.asyncio + async def test_find_latest_for_key_reads_the_keys_own_partition_newest_router_first(self): + repo, lookups = self._repo(dict(self.ROW)) + row = await repo.find_latest_for_key("hashed-key", "s1") + assert lookups == [{"where": {"api_key": "hashed-key", "session_id": "s1"}, "order": {"last_turn_at": "desc"}}] + assert row is not None + assert (row.router_name, row.turns, row.spend, row.saved_spend) == ("claude-auto", 3, 0.14, 0.24) + assert row.baseline_models == {"anthropic/claude-opus-5": 3} + assert row.baseline_model == "anthropic/claude-opus-5" + + @pytest.mark.asyncio + async def test_find_latest_for_key_is_none_when_the_key_wrote_no_such_session(self): + repo, _ = self._repo(None) + assert await repo.find_latest_for_key("hashed-key", "unknown") is None + + def test_table_is_the_session_rollup_and_needs_a_database(self): + from litellm.repositories.autorouter_session_repository import AutoRouterSessionRepository + + client = MagicMock() + assert AutoRouterSessionRepository(client).table is client.db.litellm_autoroutersession + with pytest.raises(RuntimeError, match="No DB Connected"): + _ = AutoRouterSessionRepository(None).table diff --git a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py index f36443db1e5..d717c4e8c89 100644 --- a/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py +++ b/tests/test_litellm/router_strategy/adaptive_router/test_adaptive_router.py @@ -236,6 +236,38 @@ async def test_record_turn_attributes_satisfaction_to_previous_response_model(): assert smart_after.alpha == pytest.approx(smart_before.alpha) +@pytest.mark.asyncio +async def test_external_default_keeps_feedback_history_without_entering_bandit_pool(): + r = _make_router() + before = r._cells[(RequestType.GENERAL, "fast")] + await r.record_turn( + session_id="fallback", + model_name="fast", + request_type=RequestType.GENERAL, + turn=Turn(user_content="fix this retry bug", assistant_content="clear the cache"), + ) + await r.record_turn( + session_id="fallback", + model_name="external-default", + request_type=RequestType.GENERAL, + turn=Turn(user_content="the fix is still broken", assistant_content="keep cache entries"), + ) + assert r._cells[(RequestType.GENERAL, "fast")].beta > before.beta + await r.record_turn( + session_id="fallback", + model_name="smart", + request_type=RequestType.GENERAL, + turn=Turn( + user_content="the fix is still broken", + assistant_content="use the corrected entry", + tool_results=[{"is_error": True, "content": "failure"}], + ), + ) + assert r._feedback_contexts["fallback"].model_name == "smart" + assert all(model != "external-default" for _, model in r._cells) + assert r.config.available_models == ["fast", "smart"] + + @pytest.mark.asyncio async def test_record_turn_bounds_feedback_contexts_and_evicts_least_recent_session(): r = _make_router() diff --git a/tests/test_litellm/router_strategy/test_complexity_router.py b/tests/test_litellm/router_strategy/test_complexity_router.py index dddaee71a63..dba44d1e2e8 100644 --- a/tests/test_litellm/router_strategy/test_complexity_router.py +++ b/tests/test_litellm/router_strategy/test_complexity_router.py @@ -5,18 +5,19 @@ Tests the rule-based complexity scoring and tier assignment logic. """ import asyncio -from collections.abc import AsyncIterator import json -from copy import deepcopy -from functools import partial import logging import sys import time -from typing import Dict, Final, List +from collections.abc import AsyncIterator, Mapping +from copy import deepcopy +from functools import partial +from typing import Dict, Final, List, Literal from unittest.mock import AsyncMock, MagicMock, patch -import pytest import httpx +import pytest +import respx from pydantic import ValidationError import litellm @@ -69,6 +70,7 @@ from litellm.router_strategy.complexity_router.tier_predictor import ( from litellm.types.router import ( Deployment, LiteLLM_Params, + RouterErrors, TaggedPreRoutingStrategy, ) from litellm.types.llms.openai import ResponsesAPIResponse @@ -2671,7 +2673,8 @@ class TestEncryptedTaskClassifier: assert call["metadata"]["user_api_key_hash"] == "caller-key-hash" assert call["proxy_server_request"]["body"]["input"] == call["input"] assert call["proxy_server_request"]["originating_request_masked"] == { - "input": [task], "metadata": {"authorization": "REDACTED"}, + "input": [task], + "metadata": {"authorization": "REDACTED"}, } assert "source-secret" not in json.dumps(call) assert "originating_request_masked" not in call["proxy_server_request"]["body"] @@ -3315,19 +3318,23 @@ class TestLLMClassifier: ] @pytest.mark.asyncio - @pytest.mark.parametrize("source_body", [ - {"model": "router", "messages": [{"role": "user", "content": "source-only"}]}, - {"model": "router", "system": "source-only", "messages": [{"role": "user", "content": "ask"}]}, - {"model": "router", "instructions": "source-only", "input": "ask"}, - ]) + @pytest.mark.parametrize( + "source_body", + [ + {"model": "router", "messages": [{"role": "user", "content": "source-only"}]}, + {"model": "router", "system": "source-only", "messages": [{"role": "user", "content": "ask"}]}, + {"model": "router", "instructions": "source-only", "input": "ask"}, + ], + ) async def test_classifier_source_is_masked_and_separate_from_provider_input( self, llm_complexity_router, mock_router_instance, source_body ): mock_router_instance.acompletion = AsyncMock(return_value=_llm_response('{"tier": "SIMPLE"}')) outcome = await llm_complexity_router.aclassify( - "classify-this-ask", request_kwargs={"proxy_server_request": { - "body": {**source_body, "metadata": {"authorization": "source-secret"}} - }} + "classify-this-ask", + request_kwargs={ + "proxy_server_request": {"body": {**source_body, "metadata": {"authorization": "source-secret"}}} + }, ) assert outcome.cause == "llm_classifier" call_kwargs = mock_router_instance.acompletion.call_args.kwargs @@ -8138,7 +8145,9 @@ class TestContextAwareClassifier: ), ), ) - def test_only_text_reminder_tails_are_ignored_for_new_asks(self, tail: list[dict[str, object]], expected: bool) -> None: + def test_only_text_reminder_tails_are_ignored_for_new_asks( + self, tail: list[dict[str, object]], expected: bool + ) -> None: from litellm.router_strategy.complexity_router.complexity_router import ( _CODEX_REMINDER_MARKERS, _newest_turn_is_human_ask, @@ -13222,6 +13231,528 @@ class TestModalityRouting: assert cache.async_set_cache.await_args.kwargs["value"] == {"model": "text-cheap", "tier": "SIMPLE"} +@pytest.mark.usefixtures("local_model_cost_map") +class TestHealthFallbackDispatch: + @pytest.fixture(autouse=True) + def httpx_transport(self, monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + + @staticmethod + def _router( + surface: str = "chat", + *, + peer: bool = False, + session: bool = False, + tagged: bool = False, + budgeted: bool = False, + config: Mapping[str, object] | None = None, + ) -> Router: + provider: Final = "anthropic/claude-sonnet-5" if surface == "messages" else "openai/gpt-5.6" + base_suffix: Final = "" if surface == "messages" else "/v1" + return Router( + model_list=[ + { + "model_name": "health-router", + "litellm_params": { + "model": "auto_router/complexity_router", + "complexity_router_default_model": (config or {}).get("default_model", "fallback"), + "complexity_router_config": { + "tiers": {"SIMPLE": ["primary", "peer"] if peer else "primary", "MEDIUM": "primary"}, + "session_affinity": session, + "deployment_affinity": False, + "max_tokens_from_tier_model": False, + **(config or {}), + }, + }, + }, + *[ + { + "model_name": name, + "litellm_params": { + "model": provider, + "api_key": "test-only", + "api_base": f"https://{name}.test{base_suffix}", + **({"tags": [name]} if tagged else {}), + **( + {"max_budget": 1.0, "budget_duration": "1d"} + if budgeted and name == "primary" + else {} + ), + }, + "model_info": {"id": f"{name}-id"}, + } + for name in ("primary", "peer", "fallback") + ], + ], + num_retries=0, + enable_health_check_routing=True, + enable_tag_filtering=tagged, + ) + + @staticmethod + def _unavailable(router: Router, model_id: str, source: Literal["health", "cooldown"]) -> None: + if source == "health": + router.health_state_cache.set_deployment_health_states( + {model_id: {"is_healthy": False, "timestamp": time.time()}} + ) + else: + router.cooldown_cache.add_deployment_to_cooldown( + model_id=model_id, + original_exception=RuntimeError("unavailable"), + exception_status=503, + cooldown_time=60, + ) + + @staticmethod + def _http_response(request: httpx.Request) -> httpx.Response: + body: Final = json.loads(request.content) + text: Final = request.url.host.split(".")[0] + payload: Final[Mapping[str, object]] + events: Final[tuple[Mapping[str, object], ...]] + if request.url.path.endswith("/responses"): + from litellm.responses.main import mock_responses_api_response + + payload = mock_responses_api_response(text).model_dump() + events = ( + {"type": "response.created", "response": {**payload, "status": "in_progress"}, "sequence_number": 0}, + { + "type": "response.output_text.delta", + "delta": text, + "item_id": "msg_test", + "output_index": 0, + "content_index": 0, + "sequence_number": 1, + }, + {"type": "response.completed", "response": payload, "sequence_number": 2}, + ) + elif request.url.path.endswith("/messages"): + payload = { + "id": "msg_test", + "type": "message", + "role": "assistant", + "model": body["model"], + "content": [{"type": "text", "text": text}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 10, "output_tokens": 1}, + } + events = ( + {"type": "message_start", "message": {**payload, "content": [], "stop_reason": None}}, + {"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}, + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": text}}, + {"type": "content_block_stop", "index": 0}, + {"type": "message_delta", "delta": {"stop_reason": "end_turn"}, "usage": {"output_tokens": 1}}, + {"type": "message_stop"}, + ) + else: + payload = { + "id": "chatcmpl-test", + "object": "chat.completion", + "created": 1, + "model": body["model"], + "choices": [{"index": 0, "message": {"role": "assistant", "content": text}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 1, "total_tokens": 11}, + } + events = ( + { + **payload, + "object": "chat.completion.chunk", + "choices": [{"index": 0, "delta": {"content": text}, "finish_reason": None}], + }, + { + **payload, + "object": "chat.completion.chunk", + "choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}], + }, + ) + if not body.get("stream"): + return httpx.Response(200, json=payload) + wire: Final = "".join( + (f"event: {event['type']}\n" if "type" in event else "") + f"data: {json.dumps(event)}\n\n" + for event in events + ) + return httpx.Response( + 200, + text=wire + ("data: [DONE]\n\n" if "type" not in events[0] else ""), + headers={"content-type": "text/event-stream"}, + ) + + @staticmethod + async def _request(router: Router, surface: str, stream: bool, metadata: dict[str, object]) -> str: + if surface == "responses": + result = await router.aresponses( + model="health-router", input="Hello!", stream=stream, litellm_metadata=metadata + ) + elif surface == "messages": + result = await router.aanthropic_messages( + model="health-router", + messages=[{"role": "user", "content": "Hello!"}], + max_tokens=32, + stream=stream, + litellm_metadata=metadata, + ) + else: + result = await router.acompletion( + model="health-router", + messages=[{"role": "user", "content": "Hello!"}], + stream=stream, + metadata=metadata, + ) + if not stream: + payload = result if isinstance(result, dict) else result.model_dump() + if surface == "responses": + return payload["output"][0]["content"][0]["text"] + if surface == "messages": + return payload["content"][0]["text"] + return payload["choices"][0]["message"]["content"] + if surface == "messages": + wire: Final = b"".join([chunk async for chunk in result]).decode() + events = tuple(json.loads(line[6:]) for line in wire.splitlines() if line.startswith("data: ")) + assert events[-1]["type"] == "message_stop" + return "".join(c["delta"]["text"] for c in events if c["type"] == "content_block_delta") + chunks: Final = [chunk.model_dump() async for chunk in result] + if surface == "responses": + assert chunks[-1]["type"] == "response.completed" + return "".join(c["delta"] for c in chunks if c["type"] == "response.output_text.delta") + assert chunks[-1]["choices"][0]["finish_reason"] == "stop" + return "".join(c["choices"][0]["delta"].get("content") or "" for c in chunks if c["choices"]) + + @pytest.mark.asyncio + @pytest.mark.parametrize("surface", ["chat", "responses", "messages"]) + @pytest.mark.parametrize("stream", [False, True]) + @pytest.mark.parametrize("source", ["health", "cooldown"]) + async def test_public_call_falls_back_and_recovers( + self, surface: str, stream: bool, source: Literal["health", "cooldown"] + ) -> None: + router: Final = self._router(surface, session=True) + self._unavailable(router, "primary-id", source) + metadata: Final[dict[str, object]] = {"session_id": "outage"} + with respx.mock(assert_all_mocked=True) as upstream: + upstream.post(host__regex=r"^(primary|peer|fallback)\.test$").mock(side_effect=self._http_response) + assert await self._request(router, surface, stream, metadata) == "fallback" + assert metadata["routing_decision"]["cause"] == "health_default_fallback" + assert "tier" not in metadata["routing_decision"] + assert "health_displaced:primary" in metadata["routing_decision"]["signals"] + assert [c.request.url.host for c in upstream.calls] == ["fallback.test"] + strategy: Final = router.complexity_routers["health-router"][0].strategy + key: Final = strategy._get_session_affinity_cache_key("outage", {}) + assert await router.cache.async_get_cache(key=key) is None + if source == "health": + router.health_state_cache.set_deployment_health_states( + {"primary-id": {"is_healthy": True, "timestamp": time.time()}} + ) + else: + router.cooldown_cache.cooldown_store.delete_cache( + router.cooldown_cache.get_cooldown_cache_key("primary-id") + ) + recovered: Final[dict[str, object]] = {"session_id": "outage"} + assert await self._request(router, surface, stream, recovered) == "primary" + assert recovered["routing_decision"]["routed_model"] == "primary" + assert [c.request.url.host for c in upstream.calls] == ["fallback.test", "primary.test"] + + @pytest.mark.asyncio + @pytest.mark.parametrize("source", ["health", "cooldown"]) + async def test_partial_group_then_peer_then_default(self, source: Literal["health", "cooldown"]) -> None: + router: Final = self._router(peer=True, session=True) + router.add_deployment( + Deployment( + model_name="primary", + litellm_params=LiteLLM_Params( + model="openai/gpt-5.6", api_key="test-only", api_base="https://primary.test/v1" + ), + model_info={"id": "primary-sibling-id"}, + ) + ) + strategy: Final = router.complexity_routers["health-router"][0].strategy + key: Final = strategy._get_session_affinity_cache_key("precedence", {}) + await router.cache.async_set_cache(key=key, value={"model": "primary", "tier": "SIMPLE"}, ttl=600) + with respx.mock(assert_all_mocked=True) as upstream: + upstream.post(host__regex=r"^(primary|peer|fallback)\.test$").mock(side_effect=self._http_response) + for model_id, expected, cause in ( + ("primary-id", "primary", "session_affinity_pin"), + ("primary-sibling-id", "peer", "health_failover"), + ("peer-id", "fallback", "health_default_fallback"), + ): + self._unavailable(router, model_id, source) + metadata: Final[dict[str, object]] = {"session_id": "precedence"} + assert await self._request(router, "chat", False, metadata) == expected + assert metadata["routing_decision"]["cause"] == cause + assert await router.cache.async_get_cache(key=key) == {"model": "primary", "tier": "SIMPLE"} + assert [c.request.url.host for c in upstream.calls] == ["primary.test", "peer.test", "fallback.test"] + + @pytest.mark.asyncio + async def test_spent_deployment_budget_falls_back_to_the_default(self, monkeypatch: pytest.MonkeyPatch) -> None: + """A spent budget leaves the tier with nothing that may serve the request, and the budget + filter reports that as a bare ValueError instead of a typed router error. Reading it as + capacity skips the recovery and fails the request the recovery exists for.""" + + async def _no_sync(*args: object, **kwargs: object) -> None: + return None + + monkeypatch.setattr( + "litellm.router_strategy.budget_limiter.RouterBudgetLimiting.periodic_sync_in_memory_spend_with_redis", + _no_sync, + ) + monkeypatch.setattr(litellm, "callbacks", []) + router: Final = self._router(budgeted=True) + limiter: Final = router.router_budget_logger + assert limiter is not None, "a deployment max_budget must install the budget limiter" + await router.cache.async_set_cache(key="deployment_spend:primary-id:1d", value=2.0) + with respx.mock(assert_all_mocked=True) as upstream: + upstream.post(host__regex=r"^(primary|fallback)\.test$").mock(side_effect=self._http_response) + metadata: Final[dict[str, object]] = {} + assert await self._request(router, "chat", False, metadata) == "fallback" + assert metadata["routing_decision"]["cause"] == "health_default_fallback" + assert [c.request.url.host for c in upstream.calls] == ["fallback.test"] + + @pytest.mark.asyncio + async def test_concurrent_tag_scopes_keep_fallbacks_request_local(self) -> None: + router: Final = self._router(tagged=True) + router.add_deployment( + Deployment( + model_name="fallback", + litellm_params=LiteLLM_Params( + model="openai/gpt-5.6", api_key="test-only", api_base="https://peer.test/v1", tags=["peer"] + ), + model_info={"id": "fallback-peer-id"}, + ) + ) + self._unavailable(router, "primary-id", "cooldown") + with respx.mock(assert_all_mocked=True) as upstream: + upstream.post(host__regex=r"^(peer|fallback)\.test$").mock(side_effect=self._http_response) + scopes: Final = tuple({"tags": [name], "session_id": name} for name in ("peer", "fallback")) + results: Final = await asyncio.gather( + *(self._request(router, "chat", False, metadata) for metadata in scopes) + ) + assert results == ["peer", "fallback"] + assert [m["tags"] for m in scopes] == [["peer"], ["fallback"]] + assert [m["routing_decision"]["routed_model"] for m in scopes] == ["fallback", "fallback"] + assert sorted(c.request.url.host for c in upstream.calls) == ["fallback.test", "peer.test"] + + @pytest.mark.asyncio + async def test_probe_preserves_consumed_request_exclusions(self) -> None: + router: Final = self._router() + self._unavailable(router, "primary-id", "cooldown") + kwargs: Final = {"_excluded_deployment_ids": ["fallback-id"], "_target_order": 1} + strategy: Final = router.complexity_routers["health-router"][0].strategy + response: Final = await strategy.async_pre_routing_hook( + model="health-router", messages=[{"role": "user", "content": "Hello!"}], request_kwargs=kwargs + ) + assert response.model == "primary" + assert kwargs == {"_excluded_deployment_ids": ["fallback-id"], "_target_order": 1} + + @pytest.mark.asyncio + @pytest.mark.parametrize("default_state", ["cooldown", "unconfigured", "same-model"]) + async def test_unavailable_default_preserves_no_deployment_error(self, default_state: str) -> None: + from litellm.types.router import RouterRateLimitError + + router: Final = self._router(config={"default_model": "primary"} if default_state == "same-model" else None) + self._unavailable(router, "primary-id", "cooldown") + if default_state == "unconfigured": + router.delete_deployment(id="fallback-id") + elif default_state == "cooldown": + self._unavailable(router, "fallback-id", "cooldown") + with respx.mock(assert_all_mocked=True) as upstream: + with pytest.raises(RouterRateLimitError, match="No deployments available"): + await self._request(router, "chat", False, {}) + assert not upstream.calls + + @pytest.mark.asyncio + @pytest.mark.parametrize("plan_active", [False, True]) + async def test_plan_floor_outage_cannot_use_untiered_default(self, plan_active: bool) -> None: + from litellm.types.router import RouterRateLimitError + + router: Final = self._router( + config={"tiers": {"SIMPLE": "primary", "MEDIUM": "peer"}, "plan_mode_min_tier": "MEDIUM"} + ) + self._unavailable(router, "primary-id", "cooldown") + self._unavailable(router, "peer-id", "cooldown") + metadata: Final = {} + with respx.mock(assert_all_mocked=True, assert_all_called=False) as upstream: + upstream.post(host="fallback.test").mock(side_effect=self._http_response) + if plan_active: + with pytest.raises(RouterRateLimitError, match="No deployments available"): + await router.acompletion( + model="health-router", + messages=[ + {"role": "system", "content": "Plan mode is active"}, + {"role": "user", "content": "Hello!"}, + ], + metadata=metadata, + ) + assert not upstream.calls + assert metadata["routing_decision"]["routed_model"] == "peer" + assert metadata["routing_decision"]["tier"] == "MEDIUM" + else: + assert await self._request(router, "chat", False, metadata) == "fallback" + + @pytest.mark.asyncio + async def test_default_dispatch_drops_displaced_tier_params(self) -> None: + router: Final = self._router( + config={"tiers": {"SIMPLE": {"model_name": "primary", "litellm_params": {"max_tokens": 9}}}} + ) + with respx.mock(assert_all_mocked=True) as upstream: + upstream.post(host__regex=r"^(primary|fallback)\.test$").mock(side_effect=self._http_response) + await router.acompletion( + model="health-router", messages=[{"role": "user", "content": "Hello!"}], max_tokens=32 + ) + assert json.loads(upstream.calls[-1].request.content)["max_completion_tokens"] == 9 + self._unavailable(router, "primary-id", "cooldown") + await router.acompletion( + model="health-router", messages=[{"role": "user", "content": "Hello!"}], max_tokens=32 + ) + assert json.loads(upstream.calls[-1].request.content)["max_completion_tokens"] == 32 + assert upstream.calls[-1].request.url.host == "fallback.test" + + @pytest.mark.asyncio + @pytest.mark.parametrize("source", ["health", "cooldown"]) + async def test_pinned_session_returns_to_primary_after_outage(self, source: Literal["health", "cooldown"]) -> None: + router: Final = self._router(session=True) + with respx.mock(assert_all_mocked=True) as upstream: + upstream.post(host__regex=r"^(primary|fallback)\.test$").mock(side_effect=self._http_response) + assert await self._request(router, "chat", False, {"session_id": "pinned"}) == "primary" + self._unavailable(router, "primary-id", source) + outage: Final[dict[str, object]] = {"session_id": "pinned"} + assert await self._request(router, "chat", False, outage) == "fallback" + assert outage["routing_decision"]["cause"] == "health_default_fallback" + if source == "health": + router.health_state_cache.set_deployment_health_states( + {"primary-id": {"is_healthy": True, "timestamp": time.time()}} + ) + else: + router.cooldown_cache.cooldown_store.delete_cache( + router.cooldown_cache.get_cooldown_cache_key("primary-id") + ) + recovered: Final[dict[str, object]] = {"session_id": "pinned"} + assert await self._request(router, "chat", False, recovered) == "primary" + assert recovered["routing_decision"]["cause"] == "session_affinity_pin" + assert [c.request.url.host for c in upstream.calls] == ["primary.test", "fallback.test", "primary.test"] + + @pytest.mark.asyncio + async def test_policy_plugin_does_not_escape_to_live_default(self) -> None: + from litellm.types.router import RouterRateLimitError, RoutingContext + + class PrimaryOnly: + async def run(self, context: RoutingContext) -> RoutingContext: + context.candidate_models = [name for name in context.candidate_models if name == "primary"] + return context + + router: Final = self._router(peer=True, config={"plugins": [PrimaryOnly()]}) + self._unavailable(router, "primary-id", "cooldown") + with respx.mock(assert_all_mocked=True) as upstream: + with pytest.raises(RouterRateLimitError, match="No deployments available"): + await self._request(router, "chat", False, {}) + assert not upstream.calls + + @pytest.mark.asyncio + @pytest.mark.parametrize("live_tier", [True, False]) + @pytest.mark.parametrize("default_fits", [True, False]) + async def test_context_recovery_precedes_default_with_prechecks_off( + self, live_tier: bool, default_fits: bool + ) -> None: + from litellm.types.router import RouterRateLimitError + + router: Final = self._router(config={"tiers": {"SIMPLE": "primary", "MEDIUM": "peer", "COMPLEX": "large"}}) + router.add_deployment( + Deployment( + model_name="large", + litellm_params=LiteLLM_Params( + model="openai/gpt-5.6", api_key="test-only", api_base="https://large.test/v1" + ), + model_info={"id": "large-id", "max_input_tokens": 10000}, + ) + ) + for deployment in router.model_list: + deployment["model_info"]["max_input_tokens"] = ( + 10 + if deployment["model_name"] == "primary" + or (deployment["model_name"] == "fallback" and not default_fits) + else 10000 + ) + self._unavailable(router, "peer-id", "cooldown") + if not live_tier: + self._unavailable(router, "large-id", "cooldown") + assert router.enable_pre_call_checks is False + metadata: Final = {} + messages: Final = [{"role": "user", "content": "hello " * 100}] + with respx.mock(assert_all_mocked=True, assert_all_called=False) as upstream: + upstream.post(host__regex=r"^(large|fallback)\.test$").mock(side_effect=self._http_response) + if not live_tier and not default_fits: + with pytest.raises(RouterRateLimitError, match="No deployments available"): + await router.acompletion(model="health-router", messages=messages, metadata=metadata) + assert not upstream.calls + else: + result: Final = await router.acompletion(model="health-router", messages=messages, metadata=metadata) + expected: Final = "large" if live_tier else "fallback" + assert result.choices[0].message.content == expected + assert upstream.calls[-1].request.url.host == f"{expected}.test" + assert metadata["routing_decision"].get("tier") == ("COMPLEX" if live_tier else None) + + @pytest.mark.asyncio + @pytest.mark.parametrize("live_tier", [True, False]) + async def test_modality_recovery_precedes_default(self, live_tier: bool) -> None: + router: Final = self._router( + config={"modality_routing": True, "tiers": {"SIMPLE": "primary", "MEDIUM": "peer", "COMPLEX": "vision"}} + ) + router.add_deployment( + Deployment( + model_name="vision", + litellm_params=LiteLLM_Params( + model="openai/gpt-5.6", api_key="test-only", api_base="https://vision.test/v1" + ), + model_info={"id": "vision-id", "supports_vision": True}, + ) + ) + for deployment in router.model_list: + deployment["model_info"]["supports_vision"] = deployment["model_name"] != "primary" + self._unavailable(router, "peer-id", "cooldown") + if not live_tier: + self._unavailable(router, "vision-id", "cooldown") + with respx.mock(assert_all_mocked=True) as upstream: + upstream.post(host__regex=r"^(vision|fallback)\.test$").mock(side_effect=self._http_response) + result: Final = await router.acompletion( + model="health-router", + messages=[ + { + "role": "user", + "content": [ + {"type": "text", "text": "Hello!"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,aGk="}}, + ], + } + ], + ) + expected: Final = "vision" if live_tier else "fallback" + assert result.choices[0].message.content == expected + assert upstream.calls[-1].request.url.host == f"{expected}.test" + + @pytest.mark.asyncio + @pytest.mark.parametrize("default_fits", [True, False]) + async def test_modality_default_must_also_fit_context(self, default_fits: bool) -> None: + router: Final = self._router(config={"modality_routing": True, "tiers": {"SIMPLE": "primary"}}) + for deployment in router.model_list: + deployment["model_info"]["supports_vision"] = deployment["model_name"] == "fallback" + deployment["model_info"]["max_input_tokens"] = 10000 if default_fits else 10 + with respx.mock(assert_all_mocked=True, assert_all_called=False) as upstream: + upstream.post(host="fallback.test").mock(side_effect=self._http_response) + messages: Final = [ + { + "role": "user", + "content": [ + {"type": "text", "text": "hello " * 100}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,aGk="}}, + ], + } + ] + if default_fits: + result: Final = await router.acompletion(model="health-router", messages=messages) + assert result.choices[0].message.content == "fallback" + else: + with pytest.raises(litellm.BadRequestError, match="modality_routing is enabled"): + await router.acompletion(model="health-router", messages=messages) + assert not upstream.calls + + class TestTierHealthFailover: """A tier whose decided model group is entirely in cooldown falls back to a live peer.""" @@ -13256,7 +13787,7 @@ class TestTierHealthFailover: probed_prompts = [] async def get_healthy_deployments( - model, request_kwargs, messages=None, input=None, parent_otel_span=None, **kwargs + model, request_kwargs, messages=None, input=None, parent_otel_span=None, health_check_probe=False ): probed_kwargs.append(request_kwargs) probed_prompts.append((messages, input)) @@ -13758,6 +14289,51 @@ class TestTierHealthFailover: for _, probed_input in router.litellm_router_instance.probed_prompts ), "the eligibility probe must forward `input` to the owner" + @pytest.mark.asyncio + @pytest.mark.parametrize( + "raised, expected", + [ + (ValueError(f"{RouterErrors.no_deployments_with_tag_routing.value}. Passed model=b"), {"live-c"}), + ( + ValueError(f"{RouterErrors.no_deployments_with_provider_budget_routing.value}: b over budget"), + {"live-c"}, + ), + (ValueError("cannot unpack non-sequence"), {"exhausted-b", "live-c"}), + ], + ) + async def test_a_marked_exhaustion_value_error_is_a_verdict_and_an_unmarked_one_is_not( + self, mock_router_instance, raised, expected + ): + """Budget and tag filters exhaust a group without a typed error, signalling it only by a + RouterErrors marker on a bare ValueError. Those are verdicts; any other ValueError is a + fault, and a fault must still read as capacity rather than silently rerouting.""" + router = self._router( + mock_router_instance, + { + "tiers": { + "SIMPLE": ["dead-a", "exhausted-b", "live-c"], + "MEDIUM": "mid", + "COMPLEX": "big", + "REASONING": "top", + }, + "session_affinity": True, + }, + {"dead-a": ["id-a1"], "exhausted-b": ["id-b1"], "live-c": ["id-c1"]}, + cooling=("id-a1",), + raises_for={"exhausted-b": raised}, + ) + key = router._get_session_affinity_cache_key("sess-exhausted", {}) + await router.litellm_router_instance.cache.async_set_cache( + key=key, value={"model": "dead-a", "tier": "SIMPLE"}, ttl=600 + ) + results = [ + await router.async_pre_routing_hook( + model="m", request_kwargs={"metadata": {"session_id": "sess-exhausted"}}, messages=self.SIMPLE_MESSAGE + ) + for _ in range(20) + ] + assert {r.model for r in results} == expected + @pytest.mark.asyncio async def test_a_group_the_router_has_no_deployment_for_is_not_a_failover_target(self, mock_router_instance): """The owner answers an unconfigured group with BadRequestError. Reading that as live diff --git a/tests/test_litellm/rust_bridge/native_route_wheel_test.py b/tests/test_litellm/rust_bridge/native_route_wheel_test.py index a7f50a82a99..8d83f4ca8a6 100644 --- a/tests/test_litellm/rust_bridge/native_route_wheel_test.py +++ b/tests/test_litellm/rust_bridge/native_route_wheel_test.py @@ -73,7 +73,7 @@ def assert_native_request( headers: HTTPMessage, body: object, ) -> None: - if route not in {"ocr", "transcription", "messages", "chat_completions"}: + if route not in {"ocr", "azure_ocr", "azure_di", "transcription", "messages", "chat_completions"}: raise AssertionError(f"unexpected route marker: {route!r}") if outcome not in {"success", "429", "hang"}: raise AssertionError(f"unexpected outcome marker: {outcome!r}") @@ -86,6 +86,19 @@ def assert_native_request( assert body["document"]["document_url"] == "https://example.com/document.pdf" assert body["include_image_base64"] is True return + if route == "azure_ocr": + assert path == "/providers/mistral/azure/ocr" + assert headers.get("authorization") == "Bearer prepared-azure-token" + assert body["model"] == "mistral-ocr-2505" + assert body["document"]["document_url"] == "data:application/pdf;base64,YWJj" + return + if route == "azure_di": + assert path.startswith("/documentintelligence/documentModels/prebuilt-read:analyze?") + assert "api-version=2024-11-30" in path + assert "pages=1%2C3" in path + assert headers.get("ocp-apim-subscription-key") == "di-key" + assert body == {"base64Source": "YWJj"} + return if route == "transcription": assert path == "/model/mistral.voxtral-mini-3b-2507/converse" assert headers.get("authorization", "").startswith("AWS4-HMAC-SHA256 ") @@ -107,8 +120,10 @@ def assert_native_request( def native_response(status: int, route: str | None) -> bytes: if status == 429: return b'{"error":"native-rate-limit"}' - if route == "ocr": + if route in {"ocr", "azure_ocr"}: return b'{"pages":[{"index":0,"markdown":"native-ocr"}]}' + if route == "azure_di": + return b'{"status":"succeeded","analyzeResult":{"pages":[]}}' if route == "transcription": return b'{"output":{"message":{"content":[{"text":"native-transcription"}]}}}' return ANTHROPIC_RESPONSE @@ -181,6 +196,32 @@ def assert_success(route: str, response: object) -> None: raise AssertionError(f"{route} returned {actual!r}, expected {expected!r}") +def azure_ocr_kwargs(api_base: str) -> dict[str, object]: + return { + "model": "mistral-ocr-2505", + "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + "api_base": api_base, + "custom_llm_provider": "azure_ai", + "extra_headers": { + "x-test-outcome": "success", + "x-test-route": "azure_ocr", + }, + "optional_params": {"azure_ad_token": "prepared-azure-token"}, + } + + +def azure_di_kwargs(api_base: str) -> dict[str, object]: + return { + "model": "doc-intelligence/prebuilt-read", + "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + "api_key": "di-key", + "api_base": api_base, + "custom_llm_provider": "azure_ai", + "extra_headers": {"x-test-outcome": "success", "x-test-route": "azure_di"}, + "optional_params": {"req_format": "native", "pages": [0, 2]}, + } + + def success_value(route: str, response: dict[object, object]) -> object: if route == "ocr": return response["pages"][0]["markdown"] @@ -211,6 +252,9 @@ def exercise_sync(native: object, api_base: str) -> None: assert_rate_limit(native, route, error) else: raise AssertionError(f"{route} accepted a 429 response") + assert_success("ocr", native.ocr(**azure_ocr_kwargs(api_base))) + di_response: Final = native.ocr(**azure_di_kwargs(api_base)) + assert di_response["provider_native_response"]["status"] == "succeeded" async def exercise_async(native: object, api_base: str) -> None: @@ -223,16 +267,14 @@ async def exercise_async(native: object, api_base: str) -> None: assert_rate_limit(native, route, error) else: raise AssertionError(f"a{route} accepted a 429 response") + assert_success("ocr", await native.aocr(**azure_ocr_kwargs(api_base))) + di_response: Final = await native.aocr(**azure_di_kwargs(api_base)) + assert di_response["provider_native_response"]["status"] == "succeeded" async def exercise_async_concurrency(native: object, api_base: str) -> None: responses: Final = await asyncio.wait_for( - asyncio.gather( - *( - native.amessages(**route_kwargs("messages", api_base, "success")) - for _ in range(32) - ) - ), + asyncio.gather(*(native.amessages(**route_kwargs("messages", api_base, "success")) for _ in range(32))), timeout=15, ) for response in responses: diff --git a/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py b/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py index 7e655b70756..03422d5433c 100644 --- a/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py +++ b/tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py @@ -5,11 +5,68 @@ Tests the write/read/delete cycle for JSON and simple string secrets. """ import json -from unittest.mock import AsyncMock, patch +from unittest.mock import AsyncMock, MagicMock, patch import pytest +import respx +import litellm from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 +from litellm.types.secret_managers.main import KeyManagementSettings + +_STATIC_CREDENTIALS = {"aws_access_key_id": "test-key", "aws_secret_access_key": "test-secret"} +_CMK_ARN = "arn:aws:kms:us-east-1:123456789012:key/11111111-2222-3333-4444-555555555555" + + +async def _create_secret_body_for_settings( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter, settings: KeyManagementSettings +) -> dict[str, object]: + """Boot the manager from settings the way the proxy does and return the CreateSecret body it posts to AWS.""" + monkeypatch.setattr(litellm, "secret_manager_client", None) + monkeypatch.setattr(litellm, "disable_aiohttp_transport", True) + monkeypatch.setattr(litellm, "in_memory_llm_clients_cache", None) + monkeypatch.setenv("AWS_REGION_NAME", "us-east-1") + AWSSecretsManagerV2.load_aws_secret_manager(use_aws_secret_manager=True, key_management_settings=settings) + manager = litellm.secret_manager_client + assert isinstance(manager, AWSSecretsManagerV2) + + route = respx_mock.post("https://secretsmanager.us-east-1.amazonaws.com/").respond( + json={"ARN": "arn", "Name": "litellm/test-key"} + ) + await manager.async_write_secret( + secret_name="litellm/test-key", + secret_value="sk-test-value", + optional_params=dict(_STATIC_CREDENTIALS), + ) + assert route.call_count == 1 + request = route.calls.last.request + assert request.headers["X-Amz-Target"] == "secretsmanager.CreateSecret" + return json.loads(request.content) + + +@pytest.mark.asyncio +async def test_create_secret_uses_customer_managed_kms_key_from_settings( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + body = await _create_secret_body_for_settings( + monkeypatch, + respx_mock, + KeyManagementSettings(store_virtual_keys=True, aws_region_name="us-east-1", kms_key_id=_CMK_ARN), + ) + assert body["KmsKeyId"] == _CMK_ARN + assert body["Name"] == "litellm/test-key" + assert body["SecretString"] == "sk-test-value" + + +@pytest.mark.asyncio +async def test_create_secret_omits_kms_key_id_when_not_configured( + monkeypatch: pytest.MonkeyPatch, respx_mock: respx.MockRouter +) -> None: + body = await _create_secret_body_for_settings( + monkeypatch, respx_mock, KeyManagementSettings(store_virtual_keys=True, aws_region_name="us-east-1") + ) + assert "KmsKeyId" not in body + assert body["Name"] == "litellm/test-key" @pytest.mark.asyncio diff --git a/tests/test_litellm/test_model_prices_schema.py b/tests/test_litellm/test_model_prices_schema.py index c2c22c25998..0b9dbd23097 100644 --- a/tests/test_litellm/test_model_prices_schema.py +++ b/tests/test_litellm/test_model_prices_schema.py @@ -4,6 +4,8 @@ import importlib.util import json import re from pathlib import Path +from types import MappingProxyType +from typing import Final import jsonschema import pytest @@ -217,3 +219,50 @@ def test_chat_latest_declares_the_one_effort_openai_accepts(prices: dict): with no declared levels resolves to None, which lets /model_group/info and the dashboard offer levels the upstream will 400 on.""" assert resolve_supported_reasoning_efforts(prices["chat-latest"], deployment_is_mapped=True) == ("medium",) + + +BEDROCK_OPENAI_GPT_MARKERS: Final = ("openai.gpt-5.4", "openai.gpt-5.5", "openai.gpt-5.6", "openai.gpt-6-astra") +BEDROCK_PROVIDERS: Final = frozenset(("bedrock", "bedrock_converse", "bedrock_mantle")) +BEDROCK_ROW_PREFIXES: Final = ("bedrock_mantle/", "us.", "global.") +GPT_5_4_BEDROCK_LADDER: Final = ("none", "low", "medium", "high", "xhigh") +GPT_5_6_BEDROCK_LADDER: Final = ("none", "low", "medium", "high", "xhigh", "max") +GPT_6_ASTRA_BEDROCK_LADDER: Final = ("low", "medium", "high", "xhigh", "max") +BEDROCK_OPENAI_GPT_LADDERS: Final = MappingProxyType( + { + "bedrock_mantle/openai.gpt-5.4": GPT_5_4_BEDROCK_LADDER, + "bedrock_mantle/openai.gpt-5.5": GPT_5_4_BEDROCK_LADDER, + **{ + f"{prefix}openai.gpt-5.6-{variant}": GPT_5_6_BEDROCK_LADDER + for prefix in BEDROCK_ROW_PREFIXES + for variant in ("luna", "sol", "terra") + }, + **{f"{prefix}openai.gpt-6-astra": GPT_6_ASTRA_BEDROCK_LADDER for prefix in BEDROCK_ROW_PREFIXES}, + } +) + + +@pytest.mark.parametrize( + ("name", "ladder"), tuple(BEDROCK_OPENAI_GPT_LADDERS.items()), ids=tuple(BEDROCK_OPENAI_GPT_LADDERS) +) +def test_bedrock_openai_gpt_rows_advertise_the_ladder_bedrock_accepts(prices: dict, name: str, ladder: tuple[str, ...]): + """Each ladder is the set of levels Bedrock answered 200 to for that row through the proxy on + 2026-09-11 (PR #40740): the Mantle rows go out over its Responses endpoint and the Converse rows + over inference profiles. Bedrock differs from the direct OpenAI rows in two places, gpt-5.6 and + gpt-6-astra take max there, and gpt-6-astra refuses none; minimal is refused on every row. + xhigh and max are opt-in for the resolver, so a row missing either flag silently drops that + level from every group it belongs to.""" + assert resolve_supported_reasoning_efforts(prices[name], deployment_is_mapped=True) == ladder + + +def test_every_bedrock_openai_gpt_row_advertises_xhigh(prices: dict): + """The GovCloud and gpt-5.6-cyber rows cannot be called from our account, so they carry the + family's xhigh flag rather than a measured ladder.""" + missing: Final = [ + name + for name, entry in prices.items() + if isinstance(entry, dict) + and entry.get("litellm_provider") in BEDROCK_PROVIDERS + and any(marker in name for marker in BEDROCK_OPENAI_GPT_MARKERS) + and "xhigh" not in (resolve_supported_reasoning_efforts(entry, deployment_is_mapped=True) or ()) + ] + assert missing == [] diff --git a/tests/test_litellm/test_responses_id_security.py b/tests/test_litellm/test_responses_id_security.py index d35b9563888..a6081670172 100644 --- a/tests/test_litellm/test_responses_id_security.py +++ b/tests/test_litellm/test_responses_id_security.py @@ -855,3 +855,242 @@ class TestAsyncPostCallSuccessHook: ) assert result == mock_response + + + +_FABRICATED_PROVIDER_RESPONSE_ID = "resp_fabricatedprovideridaaaaaaaaaaaaaaaa" +_FABRICATED_UNMANAGED_ID = "resp_fabricatedunmanagedidbbbbbbbbbbbbbbbb" +_UNIT_TEST_SALT_KEY = "lit6837-unit-test-salt-key" +_ADDRESSED_ID_FIELD_BY_CALL_TYPE = { + "aresponses": "previous_response_id", + "aget_responses": "response_id", + "adelete_responses": "response_id", + "acancel_responses": "response_id", + "alist_input_items": "response_id", +} + + +@pytest.fixture +def salt_key_env(monkeypatch): + """Give the encrypt/decrypt helpers a real salt key so ids round-trip for real.""" + monkeypatch.setenv("LITELLM_SALT_KEY", _UNIT_TEST_SALT_KEY) + return _UNIT_TEST_SALT_KEY + + +def _hook(general_settings=None, signing_key=_UNIT_TEST_SALT_KEY): + settings = general_settings if general_settings is not None else {} + return ResponsesIDSecurity( + general_settings_reader=lambda: settings, + signing_key_reader=lambda: signing_key, + ) + + +def _auth(user_id="owner-user", team_id="owner-team", user_role=None): + from litellm.proxy._types import UserAPIKeyAuth + + return UserAPIKeyAuth(user_id=user_id, team_id=team_id, user_role=user_role) + + +def _issue_managed_id(hook, owner, provider_response_id=_FABRICATED_PROVIDER_RESPONSE_ID): + """Mint an id exactly the way the proxy hands one to a client on create.""" + issued = hook._encrypt_response_id( + ResponsesAPIResponse( + id=provider_response_id, created_at=1234567890, output=[], status="completed" + ), + owner, + ) + return issued.id + + +class TestUnrecognizedResponseIdIsRejected: + """An id this proxy never issued carries no owner, so it must not reach the provider.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize("call_type", sorted(_ADDRESSED_ID_FIELD_BY_CALL_TYPE)) + async def test_unmanaged_id_is_rejected_and_not_forwarded(self, mock_cache, salt_key_env, call_type): + field = _ADDRESSED_ID_FIELD_BY_CALL_TYPE[call_type] + data = {field: _FABRICATED_UNMANAGED_ID} + + with pytest.raises(HTTPException) as exc_info: + await _hook().async_pre_call_hook( + user_api_key_dict=_auth(), + cache=mock_cache, + data=data, + call_type=call_type, + ) + + assert exc_info.value.status_code == 403 + assert "allow_unmanaged_response_ids" in exc_info.value.detail + assert data[field] == _FABRICATED_UNMANAGED_ID + + @pytest.mark.asyncio + async def test_owner_can_still_address_the_id_the_proxy_issued_it(self, mock_cache, salt_key_env): + hook = _hook() + owner = _auth() + data = {"response_id": _issue_managed_id(hook, owner)} + + result = await hook.async_pre_call_hook( + user_api_key_dict=owner, + cache=mock_cache, + data=data, + call_type="aget_responses", + ) + + assert result["response_id"] == _FABRICATED_PROVIDER_RESPONSE_ID + + @pytest.mark.asyncio + async def test_stranger_cannot_address_an_id_issued_to_someone_else(self, mock_cache, salt_key_env): + hook = _hook() + issued_id = _issue_managed_id(hook, _auth()) + data = {"response_id": issued_id} + + with pytest.raises(HTTPException) as exc_info: + await hook.async_pre_call_hook( + user_api_key_dict=_auth(user_id="stranger-user", team_id="stranger-team"), + cache=mock_cache, + data=data, + call_type="aget_responses", + ) + + assert exc_info.value.status_code == 403 + assert data["response_id"] == issued_id + + @pytest.mark.asyncio + async def test_unmanaged_previous_response_id_cannot_seed_a_new_response(self, mock_cache, salt_key_env): + data = {"model": "gpt-fake", "previous_response_id": _FABRICATED_UNMANAGED_ID} + + with pytest.raises(HTTPException) as exc_info: + await _hook().async_pre_call_hook( + user_api_key_dict=_auth(), + cache=mock_cache, + data=data, + call_type="aresponses", + ) + + assert exc_info.value.status_code == 403 + assert data["previous_response_id"] == _FABRICATED_UNMANAGED_ID + + @pytest.mark.asyncio + async def test_re_entering_the_hook_on_the_same_request_does_not_reject(self, mock_cache, salt_key_env): + """The rate-limit fallback retry runs pre-call twice over one already-rewritten dict.""" + hook = _hook() + owner = _auth() + data = {"model": "gpt-fake", "previous_response_id": _issue_managed_id(hook, owner)} + + first = await hook.async_pre_call_hook( + user_api_key_dict=owner, cache=mock_cache, data=data, call_type="aresponses" + ) + second = await hook.async_pre_call_hook( + user_api_key_dict=owner, cache=mock_cache, data=first, call_type="aresponses" + ) + + assert second["previous_response_id"] == _FABRICATED_PROVIDER_RESPONSE_ID + + +class TestUnmanagedResponseIdEscapeHatches: + """Deployments that pass provider ids through on purpose must keep working.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize( + "general_settings", + [{"allow_unmanaged_response_ids": True}, {"disable_responses_id_security": True}], + ) + async def test_opted_in_settings_forward_the_id_untouched(self, mock_cache, salt_key_env, general_settings): + data = {"response_id": _FABRICATED_UNMANAGED_ID} + + result = await _hook(general_settings=general_settings).async_pre_call_hook( + user_api_key_dict=_auth(), + cache=mock_cache, + data=data, + call_type="aget_responses", + ) + + assert result["response_id"] == _FABRICATED_UNMANAGED_ID + + @pytest.mark.asyncio + async def test_proxy_without_a_signing_key_forwards_the_id_untouched(self, mock_cache, monkeypatch): + monkeypatch.delenv("LITELLM_SALT_KEY", raising=False) + data = {"response_id": _FABRICATED_UNMANAGED_ID} + + result = await _hook(signing_key=None).async_pre_call_hook( + user_api_key_dict=_auth(), + cache=mock_cache, + data=data, + call_type="aget_responses", + ) + + assert result["response_id"] == _FABRICATED_UNMANAGED_ID + + @pytest.mark.asyncio + async def test_proxy_admin_may_address_an_unmanaged_id(self, mock_cache, salt_key_env): + from litellm.proxy._types import LitellmUserRoles + + data = {"response_id": _FABRICATED_UNMANAGED_ID} + + result = await _hook().async_pre_call_hook( + user_api_key_dict=_auth(user_role=LitellmUserRoles.PROXY_ADMIN), + cache=mock_cache, + data=data, + call_type="aget_responses", + ) + + assert result["response_id"] == _FABRICATED_UNMANAGED_ID + + +class TestClientSuppliedRetainedIdCannotBypassAuthorization: + """The retained-id key travels in the request body, so it is re-authorized, never trusted.""" + + @pytest.mark.asyncio + @pytest.mark.parametrize("call_type", sorted(_ADDRESSED_ID_FIELD_BY_CALL_TYPE)) + async def test_forged_retained_id_is_still_authorized(self, mock_cache, salt_key_env, call_type): + field = _ADDRESSED_ID_FIELD_BY_CALL_TYPE[call_type] + data = { + field: _FABRICATED_UNMANAGED_ID, + "_litellm_addressed_response_id": _FABRICATED_UNMANAGED_ID, + } + + with pytest.raises(HTTPException) as exc_info: + await _hook().async_pre_call_hook( + user_api_key_dict=_auth(), + cache=mock_cache, + data=data, + call_type=call_type, + ) + + assert exc_info.value.status_code == 403 + assert data[field] == _FABRICATED_UNMANAGED_ID + + @pytest.mark.asyncio + @pytest.mark.parametrize("forged", [{"nested": "value"}, ["list"], 42, "", None]) + async def test_non_string_retained_id_falls_back_to_the_addressed_field(self, mock_cache, salt_key_env, forged): + data = {"response_id": _FABRICATED_UNMANAGED_ID, "_litellm_addressed_response_id": forged} + + with pytest.raises(HTTPException) as exc_info: + await _hook().async_pre_call_hook( + user_api_key_dict=_auth(), + cache=mock_cache, + data=data, + call_type="aget_responses", + ) + + assert exc_info.value.status_code == 403 + + @pytest.mark.asyncio + async def test_stranger_forging_their_own_id_never_reaches_someone_elses_response( + self, mock_cache, salt_key_env + ): + hook = _hook() + stranger = _auth(user_id="stranger-user", team_id="stranger-team") + stranger_id = _issue_managed_id(hook, stranger, provider_response_id="resp_strangerownprovideridcccccccc") + victim_provider_id = "resp_victimprovideriddddddddddddddddddddd" + data = {"response_id": victim_provider_id, "_litellm_addressed_response_id": stranger_id} + + result = await hook.async_pre_call_hook( + user_api_key_dict=stranger, + cache=mock_cache, + data=data, + call_type="aget_responses", + ) + + assert result["response_id"] == "resp_strangerownprovideridcccccccc" + assert result["response_id"] != victim_provider_id diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index def67ccf88b..2a97e92396a 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -7381,6 +7381,63 @@ async def test_async_get_fully_unhealthy_model_names_marks_name_when_all_unhealt assert await router.async_get_fully_unhealthy_model_names() == {"gpt-4o"} +@pytest.mark.asyncio +@pytest.mark.parametrize("health_check_probe", [False, True]) +@pytest.mark.parametrize( + "state, health_routing, fails_policy, scoped, strict_ids", + [ + ("absent", True, False, False, ("dep-0", "dep-1")), + ("partial", True, False, False, ("dep-1",)), + ("all", True, False, False, ()), + ("stale", True, False, False, ("dep-0", "dep-1")), + ("all", False, False, False, ("dep-0", "dep-1")), + ("all", True, True, False, ("dep-0", "dep-1")), + ("all", True, True, True, ()), + ], +) +async def test_health_probe_preserves_normal_caller_policy( + health_check_probe: bool, + state: str, + health_routing: bool, + fails_policy: bool, + scoped: bool, + strict_ids: tuple[str, ...], +) -> None: + import time + from litellm.types.router import AllowedFailsPolicy, RouterRateLimitError + + router: Final = Router( + model_list=[ + { + "model_name": "health-group", + "litellm_params": {"model": "openai/gpt-5.6", "api_key": "test-only"}, + "model_info": {"id": model_id}, + } + for model_id in ("dep-0", "dep-1") + ], + enable_health_check_routing=health_routing, + allowed_fails_policy=AllowedFailsPolicy(ServiceUnavailableErrorAllowedFails=2) if fails_policy else None, + background_health_check_model_groups=["health-group"] if scoped else None, + ) + if state != "absent": + _seed_unhealthy_states( + router, + ("dep-0",) if state == "partial" else ("dep-0", "dep-1"), + time.time() - router.health_state_cache.staleness_threshold - 10 if state == "stale" else None, + ) + expected: Final = strict_ids if strict_ids or health_check_probe else ("dep-0", "dep-1") + if not expected: + with pytest.raises(RouterRateLimitError, match="No deployments available"): + await router.async_get_healthy_deployments(model="health-group", request_kwargs={}, health_check_probe=True) + else: + deployments: Final = await router.async_get_healthy_deployments( + model="health-group", request_kwargs={}, health_check_probe=health_check_probe + ) + assert {d["model_info"]["id"] for d in deployments} == set(expected) + assert await router.cooldown_cache.async_get_active_cooldowns(["dep-0", "dep-1"], parent_otel_span=None) == [] + + + @pytest.mark.asyncio async def test_async_get_fully_unhealthy_model_names_keeps_name_when_partial(): router = _router_with_two_deployments([False, False]) diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index 7753cb7d770..c7e46829aba 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -198,6 +198,8 @@ def test_get_model_info_resolves_provider_prefixed_model_ids(local_model_cost_ma ("perplexity/perplexity/kimi-k3", True), ("perplexity/perplexity/deepseek-v4-flash-0731", True), ("perplexity/perplexity/kimi-k2.7-code", False), + ("perplexity/perplexity/nemotron-3.5-lightning-30b-a3b", True), + ("perplexity/perplexity/nemotron-3-ultra-550b-a55b", True), ): assert litellm.supports_reasoning(model=model) is reasoning, model @@ -209,6 +211,18 @@ def test_get_model_info_resolves_provider_prefixed_model_ids(local_model_cost_ma assert via_provider["output_cost_per_token"] == 4.4e-06 assert via_provider["mode"] == "responses" + lightning = litellm.get_model_info( + model="perplexity/nemotron-3.5-lightning-30b-a3b", custom_llm_provider="perplexity" + ) + assert lightning["key"] == "perplexity/perplexity/nemotron-3.5-lightning-30b-a3b" + assert lightning["input_cost_per_token"] == 1.15e-08 + assert lightning["output_cost_per_token"] == 1.7e-07 + assert lightning["cache_read_input_token_cost"] == 1.15e-09 + assert lightning["mode"] == "responses" + + ultra = litellm.get_model_info(model="perplexity/perplexity/nemotron-3-ultra-550b-a55b") + assert ultra["key"] == "perplexity/perplexity/nemotron-3-ultra-550b-a55b" + def test_get_model_info_strips_openai_finetune_ids_without_a_custom_suffix(local_model_cost_map): info = litellm.get_model_info(model="ft:gpt-4o-2024-08-06:my-org::abc123", custom_llm_provider="openai") @@ -4075,7 +4089,7 @@ def test_deepseek_v4_models_in_cost_map(): configured in model_prices_and_context_window.json. Prices sourced from https://api-docs.deepseek.com/quick_start/pricing: - - deepseek-v4-flash: $0.44/M input, $1.32/M output + - deepseek-v4-flash: $0.30/M input, $1.20/M output - deepseek-v4-pro: $1.32/M input, $3.96/M output Closes https://github.com/BerriAI/litellm/issues/26709 @@ -4088,9 +4102,9 @@ def test_deepseek_v4_models_in_cost_map(): model_cost = json.load(f) # --- bare model names --- - for key, expected_input, expected_output, expected_cache in [ - ("deepseek-v4-flash", 4.4e-07, 1.32e-06, 1.4e-08), - ("deepseek-v4-pro", 1.32e-06, 3.96e-06, 4.4e-08), + for key, expected_input, expected_output, expected_cache, expected_vision in [ + ("deepseek-v4-flash", 3e-07, 1.2e-06, 6e-09, True), + ("deepseek-v4-pro", 1.32e-06, 3.96e-06, 4.4e-08, False), ]: info = model_cost.get(key) assert info is not None, f"{key} missing from model_prices_and_context_window.json" @@ -4102,11 +4116,12 @@ def test_deepseek_v4_models_in_cost_map(): assert info["max_input_tokens"] == 1_000_000 assert info["supports_function_calling"] is True assert info["supports_tool_choice"] is True + assert info.get("supports_vision", False) is expected_vision # --- provider-prefixed names --- - for key, expected_input, expected_output, expected_cache in [ - ("deepseek/deepseek-v4-flash", 4.4e-07, 1.32e-06, 1.4e-08), - ("deepseek/deepseek-v4-pro", 1.32e-06, 3.96e-06, 4.4e-08), + for key, expected_input, expected_output, expected_cache, expected_vision in [ + ("deepseek/deepseek-v4-flash", 3e-07, 1.2e-06, 6e-09, True), + ("deepseek/deepseek-v4-pro", 1.32e-06, 3.96e-06, 4.4e-08, False), ]: info = model_cost.get(key) assert info is not None, f"{key} missing from model_prices_and_context_window.json" @@ -4117,6 +4132,7 @@ def test_deepseek_v4_models_in_cost_map(): assert info["cache_read_input_token_cost"] == expected_cache assert info["supports_function_calling"] is True assert info["supports_tool_choice"] is True + assert info.get("supports_vision", False) is expected_vision def test_deepseek_v4_models_in_backup_cost_map(): @@ -4132,9 +4148,9 @@ def test_deepseek_v4_models_in_backup_cost_map(): model_cost = json.load(f) # --- bare model names --- - for key, expected_input, expected_output, expected_cache in [ - ("deepseek-v4-flash", 4.4e-07, 1.32e-06, 1.4e-08), - ("deepseek-v4-pro", 1.32e-06, 3.96e-06, 4.4e-08), + for key, expected_input, expected_output, expected_cache, expected_vision in [ + ("deepseek-v4-flash", 3e-07, 1.2e-06, 6e-09, True), + ("deepseek-v4-pro", 1.32e-06, 3.96e-06, 4.4e-08, False), ]: info = model_cost.get(key) assert info is not None, f"{key} missing from backup JSON" @@ -4144,11 +4160,12 @@ def test_deepseek_v4_models_in_backup_cost_map(): assert info["output_cost_per_token"] == expected_output assert info["cache_read_input_token_cost"] == expected_cache assert info["max_input_tokens"] == 1_000_000 + assert info.get("supports_vision", False) is expected_vision # --- provider-prefixed names --- - for key, expected_input, expected_output, expected_cache in [ - ("deepseek/deepseek-v4-flash", 4.4e-07, 1.32e-06, 1.4e-08), - ("deepseek/deepseek-v4-pro", 1.32e-06, 3.96e-06, 4.4e-08), + for key, expected_input, expected_output, expected_cache, expected_vision in [ + ("deepseek/deepseek-v4-flash", 3e-07, 1.2e-06, 6e-09, True), + ("deepseek/deepseek-v4-pro", 1.32e-06, 3.96e-06, 4.4e-08, False), ]: info = model_cost.get(key) assert info is not None, f"{key} missing from backup JSON" @@ -4157,6 +4174,43 @@ def test_deepseek_v4_models_in_backup_cost_map(): assert info["input_cost_per_token"] == expected_input assert info["output_cost_per_token"] == expected_output assert info["cache_read_input_token_cost"] == expected_cache + assert info.get("supports_vision", False) is expected_vision + + +def test_deprecation_dates_for_retired_xai_and_groq_models(): + import json + from pathlib import Path + + json_path = Path(__file__).parents[2] / "model_prices_and_context_window.json" + with open(json_path) as f: + model_cost = json.load(f) + + assert model_cost["xai/grok-imagine-image-quality"]["deprecation_date"] == "2026-11-02" + assert model_cost["xai/grok-imagine-image-quality-latest"]["deprecation_date"] == "2026-11-02" + assert model_cost["xai/grok-imagine-image-quality-20260403"]["deprecation_date"] == "2026-11-02" + assert model_cost["groq/gemma-7b-it"]["deprecation_date"] == "2024-12-18" + + +@pytest.mark.usefixtures("local_model_cost_map") +def test_deepseek_flash_completion_cost(): + from litellm.types.utils import ModelResponse + + response = ModelResponse( + model="deepseek-flash", + usage=Usage( + prompt_tokens=1_000_000, + completion_tokens=1_000_000, + total_tokens=2_000_000, + ), + ) + + cost = litellm.completion_cost( + completion_response=response, + model="deepseek-flash", + custom_llm_provider="deepseek", + ) + + assert cost == pytest.approx(1.50, abs=1e-9) _FIREWORKS_MODELS = [ diff --git a/tests/test_litellm_rust/test_ocr.py b/tests/test_litellm_rust/test_ocr.py index 1293de9ee0e..ad1c8c652bb 100644 --- a/tests/test_litellm_rust/test_ocr.py +++ b/tests/test_litellm_rust/test_ocr.py @@ -1,33 +1,48 @@ import json import threading from collections.abc import Generator -from dataclasses import dataclass from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer from typing import Final import pytest +import litellm from litellm.rust_bridge import ocr as rust_ocr_bridge pytestmark = pytest.mark.requires_rust_extension -@dataclass(frozen=True, slots=True) -class RecordedOCRRequest: - body: object - - @pytest.fixture -def ocr_server() -> Generator[tuple[ThreadingHTTPServer, list[RecordedOCRRequest]]]: - requests: Final[list[RecordedOCRRequest]] = [] +def ocr_server() -> Generator[tuple[ThreadingHTTPServer, list[dict[str, object]]]]: + requests: Final[list[dict[str, object]]] = [] class Handler(BaseHTTPRequestHandler): def do_POST(self) -> None: requests.append( - RecordedOCRRequest( - body=json.loads(self.rfile.read(int(self.headers["Content-Length"]))), - ) + { + "headers": {name.lower(): value for name, value in self.headers.items()}, + "body": json.loads(self.rfile.read(int(self.headers["Content-Length"]))), + } ) + if self.headers.get("x-test-stall") == "true": + self.connection.settimeout(2) + try: + self.rfile.read(1) + except TimeoutError: + pass + return + if self.headers.get("User-Agent", "").startswith("python-httpx"): + self.send_response(418) + self.end_headers() + return + status = int(self.headers.get("x-test-status", "200")) + if status != 200: + body = b'{"error":"provider unavailable"}' + self.send_response(status) + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + return response: Final = json.dumps( { "pages": [{"index": 0, "markdown": "native OCR response", "images": [], "dimensions": None}], @@ -56,7 +71,7 @@ def ocr_server() -> Generator[tuple[ThreadingHTTPServer, list[RecordedOCRRequest def test_native_ocr_with_compiled_rust_extension( - ocr_server: tuple[ThreadingHTTPServer, list[RecordedOCRRequest]], + ocr_server: tuple[ThreadingHTTPServer, list[dict[str, object]]], ) -> None: server, requests = ocr_server address: Final = server.server_address @@ -77,7 +92,143 @@ def test_native_ocr_with_compiled_rust_extension( assert response is not None assert response["pages"][0]["markdown"] == "native OCR response" assert len(requests) == 1 - assert requests[0].body == { + assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx") + assert requests[0]["body"] == { "model": "mistral-ocr-latest", "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, } + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.parametrize("model", ["mistral/mistral-ocr-latest", "azure_ai/doc-intelligence/prebuilt-read"]) +@pytest.mark.asyncio +async def test_native_public_ocr_matches_python(model, asynchronous): + import json + from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer + from threading import Thread + from typing import Final + from urllib.parse import parse_qsl, urlsplit + + from litellm.rust_bridge import _native + + assert callable(_native.ocr) + calls: Final = [] + + class Handler(BaseHTTPRequestHandler): + def do_POST(self): + body: Final = json.loads(self.rfile.read(int(self.headers["Content-Length"]))) + target: Final = urlsplit(self.path) + calls.append( + ( + target.path, + parse_qsl(target.query), + self.headers.get("Authorization"), + self.headers.get("Ocp-Apim-Subscription-Key"), + body, + ) + ) + payload: Final = ( + {"status": "succeeded", "analyzeResult": {"pages": []}} + if "doc-intelligence" in model + else {"pages": [{"index": 0, "markdown": "hello"}]} + ) + encoded: Final = json.dumps(payload).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(encoded))) + self.end_headers() + self.wfile.write(encoded) + + def log_message(self, *_args): + pass + + server: Final = ThreadingHTTPServer(("127.0.0.1", 0), Handler) + thread: Final = Thread(target=server.serve_forever, daemon=True) + thread.start() + responses: Final = [] + try: + for enabled in (False, True): + litellm.rust(enabled) + arguments: Final = { + "model": model, + "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + "api_key": "test-key", + "api_base": f"http://127.0.0.1:{server.server_port}", + "pages": [0, 2], + "timeout": 3.0, + } + response: Final = await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments) + responses.append(response.model_dump()) + assert len(calls) == 2 + assert calls[0] == calls[1] + for key in ("model", "pages", "object"): + assert responses[0][key] == responses[1][key] + finally: + server.shutdown() + server.server_close() + thread.join(timeout=3) + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.asyncio +async def test_native_ocr_failures_do_not_retry_on_python(ocr_server, asynchronous): + server, requests = ocr_server + arguments = { + "model": "mistral-ocr-latest", + "custom_llm_provider": "mistral", + "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + "api_key": "test-key", + "api_base": f"http://127.0.0.1:{server.server_port}", + "extra_headers": {"x-test-status": "503"}, + "num_retries": 0, + } + litellm.rust(True) + with pytest.raises(litellm.ServiceUnavailableError) as caught: + await litellm.aocr(**arguments) if asynchronous else litellm.ocr(**arguments) + assert caught.value.status_code == 503 + assert len(requests) == 1 + assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx") + + +@pytest.mark.parametrize("custom_provider", ["mistral", "not-a-provider"]) +def test_native_ocr_rejects_invalid_input_before_network(ocr_server, custom_provider): + from litellm.rust_bridge import _native + + server, requests = ocr_server + with pytest.raises(ValueError, match=r"invalid (OCR request field|provider)|invalid request"): + _native.ocr( + model="mistral-ocr-latest", + custom_llm_provider=custom_provider, + document={"type": "document_url"}, + api_key="test-key", + api_base=f"http://127.0.0.1:{server.server_port}", + ) + assert requests == [] + + +@pytest.mark.parametrize("asynchronous", [False, True]) +@pytest.mark.asyncio +async def test_native_ocr_enforces_request_deadline_without_fallback(ocr_server, asynchronous): + import asyncio + import time + + server, requests = ocr_server + litellm.rust(True) + arguments = { + "model": "mistral/mistral-ocr-latest", + "document": {"type": "document_url", "document_url": "data:application/pdf;base64,YWJj"}, + "api_key": "test-key", + "api_base": f"http://127.0.0.1:{server.server_port}", + "extra_headers": {"x-test-stall": "true"}, + "timeout": 0.1, + "num_retries": 0, + } + started = time.monotonic() + with pytest.raises(litellm.APIConnectionError): + await asyncio.wait_for( + litellm.aocr(**arguments) if asynchronous else asyncio.to_thread(litellm.ocr, **arguments), + timeout=3, + ) + assert 0.09 <= time.monotonic() - started < 3 + assert len(requests) == 1 + assert not requests[0]["headers"].get("user-agent", "").startswith("python-httpx") diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx index 295446186b1..e760d0b8dc3 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.test.tsx @@ -258,6 +258,44 @@ describe("RequestLogsPanel", () => { }); }); + it("jumps straight to the last page without a cursor when the last-page button is clicked", async () => { + const firstPage = Array.from({ length: 25 }, (_, index) => logEntry({ request_id: `req-${index}` })); + const lastPage = Array.from({ length: 10 }, (_, index) => logEntry({ request_id: `req-last-${index}` })); + vi.mocked(uiSpendLogsCall).mockImplementation(async ({ page }) => + page === 3 + ? { + data: lastPage, + total: 60, + page: 3, + page_size: 25, + total_pages: 3, + next_session_cursor: null, + has_more: false, + } + : { + data: firstPage, + total: 60, + page: 1, + page_size: 25, + total_pages: 3, + next_session_cursor: "2026-07-07 09:50:13|key-1|sess-1", + has_more: true, + }, + ); + renderPanel(); + + await waitFor(() => expect(row("req-0")).not.toBeNull()); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 1 of 3"); + fireEvent.click(screen.getByTestId("pagination-last")); + + await waitFor(() => expect(row("req-last-0")).not.toBeNull()); + expect(lastCall()?.page).toBe(3); + expect(lastCall()?.params?.session_cursor).toBeUndefined(); + expect(screen.getByTestId("pagination-page")).toHaveTextContent("Page 3 of 3"); + expect(screen.getByTestId("pagination-range")).toHaveTextContent("Showing 51-60 of 60"); + expect(vi.mocked(uiSpendLogsCall).mock.calls.filter(([options]) => options.page === 2)).toHaveLength(0); + }); + it("drops the cursor and returns to the first page when a filter changes", async () => { const firstPage = Array.from({ length: 50 }, (_, index) => logEntry({ request_id: `req-${index}` })); vi.mocked(uiSpendLogsCall).mockResolvedValue({ diff --git a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx index 6e984297bf2..4f39bb3b79b 100644 --- a/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx +++ b/ui/litellm-dashboard/src/components/view_logs/RequestLogsPanel.tsx @@ -209,15 +209,14 @@ export default function RequestLogsPanel({ accessToken, token, userRole, userID, setPagination({ ...requested, pageIndex: 0 }); return; } - if (requested.pageIndex <= pagination.pageIndex) { + if (requested.pageIndex !== pagination.pageIndex + 1) { setPagination(requested); return; } const nextCursor = filteredLogs.next_session_cursor; if (!nextCursor || logsQuery.isPlaceholderData) return; - const nextPageIndex = pagination.pageIndex + 1; - setSessionCursors((previous) => ({ ...previous, [nextPageIndex]: nextCursor })); - setPagination({ ...requested, pageIndex: nextPageIndex }); + setSessionCursors((previous) => ({ ...previous, [requested.pageIndex]: nextCursor })); + setPagination(requested); }, [usesSessionCursor, pagination, filteredLogs.next_session_cursor, logsQuery.isPlaceholderData], ); diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index 29435c31aee..dd455cba44a 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -1242,6 +1242,31 @@ export interface paths { patch?: never; trace?: never; }; + "/auto_router/session": { + parameters: { + query?: never; + header?: never; + path?: never; + cookie?: never; + }; + /** + * Get Auto Router Session + * @description One auto-routed session, for the key that ran it: the model its last turn was routed to and the + * session's spend against the router's savings baseline. Built for a coding agent's status line + * or stop hook, so any virtual key may call it and only ever sees rows written under its own + * key hash. Reads the LiteLLM_AutoRouterSession rollup, which the asynchronous spend flush + * fills a moment after each turn; a session with no flushed auto-routed turn yet is a 404. The + * id is bounded the way the writer bounded it, so an oversized client id still finds its row. + */ + get: operations["get_auto_router_session_auto_router_session_get"]; + put?: never; + post?: never; + delete?: never; + options?: never; + head?: never; + patch?: never; + trace?: never; + }; "/auto_router/shadow_eval": { parameters: { query?: never; @@ -9547,26 +9572,6 @@ export interface paths { patch?: never; trace?: never; }; - "/openai/": { - parameters: { - query?: never; - header?: never; - path?: never; - cookie?: never; - }; - /** - * WebSocket: openai_websocket_proxy_route - * @description WebSocket connection endpoint - */ - get: operations["websocket_openai_websocket_proxy_route_get"]; - put?: never; - post?: never; - delete?: never; - options?: never; - head?: never; - patch?: never; - trace?: never; - }; "/openai/deployments/{model}/chat/completions": { parameters: { query?: never; @@ -10122,26 +10127,6 @@ export interface paths { patch: operations["openai_proxy_route_openai__endpoint__patch"]; trace?: never; }; - "/openai_passthrough/": { - parameters: { - query?: never; - header?: never; - path?: never; - cookie?: never; - }; - /** - * WebSocket: openai_websocket_proxy_route - * @description WebSocket connection endpoint - */ - get: operations["websocket_openai_websocket_proxy_route_get_2"]; - put?: never; - post?: never; - delete?: never; - options?: never; - head?: never; - patch?: never; - trace?: never; - }; "/openai_passthrough/{endpoint}": { parameters: { query?: never; @@ -10150,132 +10135,72 @@ export interface paths { cookie?: never; }; /** - * Openai Proxy Route - * @description Pass-through endpoint for OpenAI API calls. - * - * Available on both routes: - * - /openai/{endpoint:path} - Standard OpenAI passthrough route - * - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API) - * - * Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts - * with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses). + * Openai Passthrough Route + * @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native + * implementations (e.g. the Responses API at /v1/responses). * * Examples: - * Standard route: - * - /openai/v1/chat/completions - * - /openai/v1/assistants - * - /openai/v1/threads - * - * Dedicated passthrough (for Responses API): * - /openai_passthrough/v1/responses * - /openai_passthrough/v1/responses/{response_id} * - /openai_passthrough/v1/responses/{response_id}/input_items * * [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough) */ - get: operations["openai_proxy_route_openai_passthrough__endpoint__get"]; + get: operations["openai_passthrough_route_openai_passthrough__endpoint__get"]; /** - * Openai Proxy Route - * @description Pass-through endpoint for OpenAI API calls. - * - * Available on both routes: - * - /openai/{endpoint:path} - Standard OpenAI passthrough route - * - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API) - * - * Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts - * with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses). + * Openai Passthrough Route + * @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native + * implementations (e.g. the Responses API at /v1/responses). * * Examples: - * Standard route: - * - /openai/v1/chat/completions - * - /openai/v1/assistants - * - /openai/v1/threads - * - * Dedicated passthrough (for Responses API): * - /openai_passthrough/v1/responses * - /openai_passthrough/v1/responses/{response_id} * - /openai_passthrough/v1/responses/{response_id}/input_items * * [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough) */ - put: operations["openai_proxy_route_openai_passthrough__endpoint__put"]; + put: operations["openai_passthrough_route_openai_passthrough__endpoint__put"]; /** - * Openai Proxy Route - * @description Pass-through endpoint for OpenAI API calls. - * - * Available on both routes: - * - /openai/{endpoint:path} - Standard OpenAI passthrough route - * - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API) - * - * Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts - * with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses). + * Openai Passthrough Route + * @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native + * implementations (e.g. the Responses API at /v1/responses). * * Examples: - * Standard route: - * - /openai/v1/chat/completions - * - /openai/v1/assistants - * - /openai/v1/threads - * - * Dedicated passthrough (for Responses API): * - /openai_passthrough/v1/responses * - /openai_passthrough/v1/responses/{response_id} * - /openai_passthrough/v1/responses/{response_id}/input_items * * [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough) */ - post: operations["openai_proxy_route_openai_passthrough__endpoint__post"]; + post: operations["openai_passthrough_route_openai_passthrough__endpoint__post"]; /** - * Openai Proxy Route - * @description Pass-through endpoint for OpenAI API calls. - * - * Available on both routes: - * - /openai/{endpoint:path} - Standard OpenAI passthrough route - * - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API) - * - * Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts - * with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses). + * Openai Passthrough Route + * @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native + * implementations (e.g. the Responses API at /v1/responses). * * Examples: - * Standard route: - * - /openai/v1/chat/completions - * - /openai/v1/assistants - * - /openai/v1/threads - * - * Dedicated passthrough (for Responses API): * - /openai_passthrough/v1/responses * - /openai_passthrough/v1/responses/{response_id} * - /openai_passthrough/v1/responses/{response_id}/input_items * * [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough) */ - delete: operations["openai_proxy_route_openai_passthrough__endpoint__delete"]; + delete: operations["openai_passthrough_route_openai_passthrough__endpoint__delete"]; options?: never; head?: never; /** - * Openai Proxy Route - * @description Pass-through endpoint for OpenAI API calls. - * - * Available on both routes: - * - /openai/{endpoint:path} - Standard OpenAI passthrough route - * - /openai_passthrough/{endpoint:path} - Dedicated passthrough route (recommended for Responses API) - * - * Use /openai_passthrough/* when you need guaranteed passthrough to OpenAI without conflicts - * with LiteLLM's native implementations (e.g., for the Responses API at /v1/responses). + * Openai Passthrough Route + * @description Dedicated pass-through to the OpenAI API with no overlap with LiteLLM's native + * implementations (e.g. the Responses API at /v1/responses). * * Examples: - * Standard route: - * - /openai/v1/chat/completions - * - /openai/v1/assistants - * - /openai/v1/threads - * - * Dedicated passthrough (for Responses API): * - /openai_passthrough/v1/responses * - /openai_passthrough/v1/responses/{response_id} * - /openai_passthrough/v1/responses/{response_id}/input_items * * [Docs](https://docs.litellm.ai/docs/pass_through/openai_passthrough) */ - patch: operations["openai_proxy_route_openai_passthrough__endpoint__patch"]; + patch: operations["openai_passthrough_route_openai_passthrough__endpoint__patch"]; trace?: never; }; "/organization/daily/activity": { @@ -21891,52 +21816,6 @@ export interface paths { patch?: never; trace?: never; }; - "/vertex-ai/{endpoint}": { - parameters: { - query?: never; - header?: never; - path?: never; - cookie?: never; - }; - /** - * Vertex Proxy Route - * @description Call LiteLLM proxy via Vertex AI SDK. - * - * [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai) - */ - get: operations["vertex_proxy_route_vertex_ai__endpoint__get_2"]; - /** - * Vertex Proxy Route - * @description Call LiteLLM proxy via Vertex AI SDK. - * - * [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai) - */ - put: operations["vertex_proxy_route_vertex_ai__endpoint__put_2"]; - /** - * Vertex Proxy Route - * @description Call LiteLLM proxy via Vertex AI SDK. - * - * [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai) - */ - post: operations["vertex_proxy_route_vertex_ai__endpoint__post_2"]; - /** - * Vertex Proxy Route - * @description Call LiteLLM proxy via Vertex AI SDK. - * - * [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai) - */ - delete: operations["vertex_proxy_route_vertex_ai__endpoint__delete_2"]; - options?: never; - head?: never; - /** - * Vertex Proxy Route - * @description Call LiteLLM proxy via Vertex AI SDK. - * - * [Docs](https://docs.litellm.ai/docs/pass_through/vertex_ai) - */ - patch: operations["vertex_proxy_route_vertex_ai__endpoint__patch_2"]; - trace?: never; - }; "/vertex_ai/discovery/{endpoint}": { parameters: { query?: never; @@ -23718,6 +23597,62 @@ export interface components { /** @description The decision record this request would have written to its log row */ routing_decision: components["schemas"]["StandardLoggingRoutingDecision"]; }; + /** + * AutoRouterSessionResponse + * @description One auto-routed session as its own key sees it: what the last turn ran on, and what the session cost + * against the router's savings baseline (the priciest model in its hardest tier). + */ + AutoRouterSessionResponse: { + /** + * Baseline Model + * @description The savings baseline most of this session's turns were priced against, recorded turn by turn, so it still names the counterfactual after the router is reconfigured or removed. None when no turn recorded one: rows from before the baseline was recorded, and adaptive and quality routers, which derive no baseline and so report no savings + */ + baseline_model: string | null; + /** + * Baseline Models + * @description Turns priced against each baseline model; more than one entry means the router's baseline changed mid-session and baseline_spend mixes both + */ + baseline_models: { + [key: string]: number; + }; + /** + * Baseline Spend + * @description spend plus saved_spend: the estimated single-model cost + */ + baseline_spend: number; + /** + * Last Model + * @description The deployment model the most recent turn was routed to + */ + last_model: string; + /** + * Router Name + * @description The auto-router alias the session's requests were sent to + */ + router_name: string; + /** + * Router Type + * @description complexity, adaptive or quality + */ + router_type: string; + /** + * Saved Spend + * @description Estimated savings against the baseline, net of classifier cost + */ + saved_spend: number; + /** Session Id */ + session_id: string; + /** + * Spend + * @description What the session's routed traffic actually cost, classifier calls included + */ + spend: number; + /** + * Turns + * @description Auto-routed turns the rollup has recorded for this session so far + */ + turns: number; + }; /** AwsSessionTag */ AwsSessionTag: { /** Key */ @@ -25775,6 +25710,11 @@ export interface components { * @description opt-in to RFC 8628 verification_uri_complete for the CLI SSO device flow, pre-filling the user_code in the browser. Off by default; intended for same-host clients where the device that starts the flow and the browser run on the same machine */ allow_cli_sso_verification_uri_complete?: boolean | null; + /** + * Allow Unmanaged Response Ids + * @description If True, lets keys address Responses API ids that this proxy did not issue (raw provider ids, or ids issued before response-id encryption was configured). Such an id carries no owner, so no ownership check can run on it; ids this proxy did issue keep full ownership enforcement. Off by default, in which case an unrecognized response id is rejected with 403 + */ + allow_unmanaged_response_ids?: boolean | null; /** * Allowed Routes * @description Proxy API Endpoints you want users to be able to access @@ -25894,6 +25834,11 @@ export interface components { * @description If True and SSO is configured (MICROSOFT_CLIENT_ID, GOOGLE_CLIENT_ID, GENERIC_CLIENT_ID, or SAML_IDP_METADATA_URL/XML), disables username/password login on /login, /v2/login, and /v3/login so SSO is the only way to reach the Admin UI. An admin locked out of the UI can still administer the proxy over the API with the master key; unset this setting and restart the proxy to restore UI username/password login. Default is False. */ disable_password_login_when_sso_enabled?: boolean | null; + /** + * Disable Responses Id Security + * @description If True, disables ownership enforcement on Responses API ids. Keys may then retrieve, cancel, delete, and chain from any response id, including ids belonging to another user or team and ids this proxy never issued. WARNING: this removes tenant isolation on /v1/responses + */ + disable_responses_id_security?: boolean | null; /** * Enable Openai Websocket Passthrough * @description Serve the OpenAI pass-through WebSocket route, which relays frames to OpenAI under the proxy's own provider credential without reading them. Off by default. @@ -26153,6 +26098,11 @@ export interface components { * @description If True and LiteLLM_SpendLogs has been converted to a range-partitioned table (db_scripts/partition_spend_logs.sql), retention cleanup drops expired partitions instead of deleting rows, and pre-creates upcoming partitions. Default is False. */ use_spend_logs_partitioning?: boolean | null; + /** + * User Api Key Cache Max Size + * @description max number of entries (virtual keys, teams, users, end users, memberships, ...) each worker keeps in its in-memory auth cache. Defaults to 200. Raise this if you have more active keys than that or auth lookups keep hitting the DB + */ + user_api_key_cache_max_size?: number | null; /** User Header Mappings */ user_header_mappings?: components["schemas"]["UserHeaderMapping"][] | null; /** @@ -36460,7 +36410,7 @@ export interface components { * Cause * @enum {string} */ - cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit"; + cause?: "heuristic_scorer" | "heuristic_v2" | "reasoning_override" | "llm_classifier" | "heuristic_first_short_circuit" | "hybrid_short_circuit" | "classifier_plugin" | "classifier_fallback" | "default_model_fallback" | "literal_keyword_match" | "semantic_keyword_match" | "plan_mode" | "housekeeping" | "modality_escalation" | "modality_pin_override" | "health_failover" | "health_default_fallback" | "session_affinity_pin" | "session_affinity_escalation" | "user_turn_continuation" | "default_fallback" | "keyword" | "quality_tier" | "bandit"; /** Classifier Cost */ classifier_cost?: number; /** Classifier Model */ @@ -41818,6 +41768,38 @@ export interface operations { }; }; }; + get_auto_router_session_auto_router_session_get: { + parameters: { + query: { + /** @description The client session id (x-*-session-id header) the turns were sent under */ + session_id: string; + }; + header?: never; + path?: never; + cookie?: never; + }; + requestBody?: never; + responses: { + /** @description Successful Response */ + 200: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["AutoRouterSessionResponse"]; + }; + }; + /** @description Validation Error */ + 422: { + headers: { + [name: string]: unknown; + }; + content: { + "application/json": components["schemas"]["HTTPValidationError"]; + }; + }; + }; + }; list_shadow_eval_jobs_auto_router_shadow_eval_get: { parameters: { query?: { @@ -52513,24 +52495,6 @@ export interface operations { }; }; }; - websocket_openai_websocket_proxy_route_get: { - parameters: { - query?: never; - header?: never; - path?: never; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description WebSocket Protocol Switched */ - 101: { - headers: { - [name: string]: unknown; - }; - content?: never; - }; - }; - }; chat_completion_openai_deployments__model__chat_completions_post: { parameters: { query?: never; @@ -53430,25 +53394,7 @@ export interface operations { }; }; }; - websocket_openai_websocket_proxy_route_get_2: { - parameters: { - query?: never; - header?: never; - path?: never; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description WebSocket Protocol Switched */ - 101: { - headers: { - [name: string]: unknown; - }; - content?: never; - }; - }; - }; - openai_proxy_route_openai_passthrough__endpoint__get: { + openai_passthrough_route_openai_passthrough__endpoint__get: { parameters: { query?: never; header?: never; @@ -53479,7 +53425,7 @@ export interface operations { }; }; }; - openai_proxy_route_openai_passthrough__endpoint__put: { + openai_passthrough_route_openai_passthrough__endpoint__put: { parameters: { query?: never; header?: never; @@ -53510,7 +53456,7 @@ export interface operations { }; }; }; - openai_proxy_route_openai_passthrough__endpoint__post: { + openai_passthrough_route_openai_passthrough__endpoint__post: { parameters: { query?: never; header?: never; @@ -53541,7 +53487,7 @@ export interface operations { }; }; }; - openai_proxy_route_openai_passthrough__endpoint__delete: { + openai_passthrough_route_openai_passthrough__endpoint__delete: { parameters: { query?: never; header?: never; @@ -53572,7 +53518,7 @@ export interface operations { }; }; }; - openai_proxy_route_openai_passthrough__endpoint__patch: { + openai_passthrough_route_openai_passthrough__endpoint__patch: { parameters: { query?: never; header?: never; @@ -67979,161 +67925,6 @@ export interface operations { }; }; }; - vertex_proxy_route_vertex_ai__endpoint__get_2: { - parameters: { - query?: never; - header?: never; - path: { - endpoint: string; - }; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description Successful Response */ - 200: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": unknown; - }; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; - }; - }; - vertex_proxy_route_vertex_ai__endpoint__put_2: { - parameters: { - query?: never; - header?: never; - path: { - endpoint: string; - }; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description Successful Response */ - 200: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": unknown; - }; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; - }; - }; - vertex_proxy_route_vertex_ai__endpoint__post_2: { - parameters: { - query?: never; - header?: never; - path: { - endpoint: string; - }; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description Successful Response */ - 200: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": unknown; - }; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; - }; - }; - vertex_proxy_route_vertex_ai__endpoint__delete_2: { - parameters: { - query?: never; - header?: never; - path: { - endpoint: string; - }; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description Successful Response */ - 200: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": unknown; - }; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; - }; - }; - vertex_proxy_route_vertex_ai__endpoint__patch_2: { - parameters: { - query?: never; - header?: never; - path: { - endpoint: string; - }; - cookie?: never; - }; - requestBody?: never; - responses: { - /** @description Successful Response */ - 200: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": unknown; - }; - }; - /** @description Validation Error */ - 422: { - headers: { - [name: string]: unknown; - }; - content: { - "application/json": components["schemas"]["HTTPValidationError"]; - }; - }; - }; - }; vertex_discovery_proxy_route_vertex_ai_discovery__endpoint__get: { parameters: { query?: never;