diff --git a/.github/workflows/create_daily_oss_branch.yml b/.github/workflows/create_daily_oss_branch.yml deleted file mode 100644 index 43de4a0e75f..00000000000 --- a/.github/workflows/create_daily_oss_branch.yml +++ /dev/null @@ -1,61 +0,0 @@ -name: Create Daily OSS Branch - -on: - schedule: - - cron: "0 16 * * 1-5" # 9am PT during daylight saving time, weekdays. - workflow_dispatch: - inputs: - date: - description: "Branch date in YYYY_MM_DD format. Defaults to today's UTC date." - required: false - type: string - -permissions: - contents: write - -jobs: - create-oss-branch: - if: github.repository == 'BerriAI/litellm' - runs-on: ubuntu-latest - timeout-minutes: 10 - - steps: - - name: Checkout repository - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - fetch-depth: 0 - persist-credentials: false - - - name: Create dated OSS branch - env: - GITHUB_TOKEN: ${{ secrets.GITHUB_TOKEN }} - REQUESTED_DATE: ${{ inputs.date }} - run: | - set -euo pipefail - - if [ -n "${REQUESTED_DATE}" ]; then - if ! echo "${REQUESTED_DATE}" | grep -Eq '^[0-9]{4}_[0-9]{2}_[0-9]{2}$'; then - echo "::error::date must use YYYY_MM_DD format, got '${REQUESTED_DATE}'" - exit 1 - fi - BRANCH_DATE="${REQUESTED_DATE}" - else - BRANCH_DATE="$(date -u +'%Y_%m_%d')" - fi - - BRANCH_NAME="litellm_oss_daily_${BRANCH_DATE}" - echo "Creating branch: ${BRANCH_NAME}" - - git config user.name "github-actions[bot]" - git config user.email "github-actions[bot]@users.noreply.github.com" - - git fetch origin main "${BRANCH_NAME}" || true - - if git show-ref --verify --quiet "refs/remotes/origin/${BRANCH_NAME}"; then - echo "Branch ${BRANCH_NAME} already exists. Skipping creation." - exit 0 - fi - - git checkout -b "${BRANCH_NAME}" origin/main - git push "https://x-access-token:${GITHUB_TOKEN}@github.com/${GITHUB_REPOSITORY}.git" "${BRANCH_NAME}" - echo "Successfully created and pushed branch: ${BRANCH_NAME}" diff --git a/.github/workflows/guard-main-branch.yml b/.github/workflows/guard-main-branch.yml index aa4968f0c1e..5bc561c6441 100644 --- a/.github/workflows/guard-main-branch.yml +++ b/.github/workflows/guard-main-branch.yml @@ -31,12 +31,12 @@ jobs: echo "PR head repo: $HEAD_REPO" echo "PR head branch: $HEAD_REF" if [ "$HEAD_REPO" != "$BASE_REPO" ]; then - echo "::error::PRs to main must originate from the canonical repository ($BASE_REPO), not a fork ($HEAD_REPO). External contributors should open PRs against the current daily OSS branch (named litellm_oss_daily_YYYY_MM_DD; a fresh one is cut each weekday, so target the most recent) instead." + echo "::error::PRs to main must originate from the canonical repository ($BASE_REPO), not a fork ($HEAD_REPO). External contributors should open PRs against 'litellm_internal_staging' instead." exit 1 fi if [ "$HEAD_REF" = "litellm_internal_staging" ] || [[ "$HEAD_REF" == litellm_hotfix_?* ]]; then echo "Allowed source branch." exit 0 fi - echo "::error::PRs to main must originate from 'litellm_internal_staging' or a 'litellm_hotfix_*' branch. Got: '$HEAD_REF'. If this is a contribution, retarget the PR against the current daily OSS branch (named litellm_oss_daily_YYYY_MM_DD; a fresh one is cut each weekday, so target the most recent) instead." + echo "::error::PRs to main must originate from 'litellm_internal_staging' or a 'litellm_hotfix_*' branch. Got: '$HEAD_REF'. If this is a contribution, retarget the PR against 'litellm_internal_staging' instead." exit 1 diff --git a/.github/workflows/oss_daily_guardrails.yml b/.github/workflows/oss_daily_guardrails.yml deleted file mode 100644 index f9dc746ee05..00000000000 --- a/.github/workflows/oss_daily_guardrails.yml +++ /dev/null @@ -1,50 +0,0 @@ -name: OSS Daily Guardrails - -on: - push: - branches: - - "litellm_oss_daily_20*" - pull_request: - branches: - - "litellm_oss_daily_20*" - - litellm_internal_staging - -permissions: - contents: read - -concurrency: - group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} - cancel-in-progress: true - -jobs: - oss-safe-checks: - name: Run OSS daily safe checks - if: startsWith(github.ref_name, 'litellm_oss_daily_20') || startsWith(github.head_ref, 'litellm_oss_daily_20') || startsWith(github.base_ref, 'litellm_oss_daily_20') - runs-on: ubuntu-latest - timeout-minutes: 10 - - steps: - - name: Checkout repository - uses: actions/checkout@08eba0b27e820071cde6df949e0beb9ba4906955 # v4.3.0 - with: - persist-credentials: false - - - name: Set up Python - uses: actions/setup-python@a26af69be951a213d495a4c3e4e4022e16d87065 # v5.6.0 - with: - python-version: "3.12" - - - name: Set up uv - uses: ./.github/actions/setup-uv-with-retries - with: - version: "0.10.9" - - - name: Run secret scan test - run: | - uv run --frozen --with 'pytest==9.0.2' pytest tests/litellm/test_no_hardcoded_secrets.py -v - - - name: Run Ruff - run: | - uv sync --frozen - cd litellm - uv run --no-sync ruff check . diff --git a/CLAUDE.md b/CLAUDE.md index 9f708716c6d..1a4826d51e9 100644 --- a/CLAUDE.md +++ b/CLAUDE.md @@ -19,7 +19,7 @@ Same thing for bug fixes. The tests should make it so that this specific bug can End-to-end tests belong in `tests/e2e/` and must follow the harness conventions documented in that directory's `CLAUDE.md` -When creating PRs, don't set base to `main`. `litellm_internal_staging` serves that purpose for internal contributors; external / OSS contributions target the current daily OSS branch instead, named `litellm_oss_daily_YYYY_MM_DD` (a fresh one is cut each weekday, so use the most recent) +When creating PRs, don't set base to `main`. `litellm_internal_staging` is the default base branch and serves that purpose for both internal and external / OSS contributions When writing a PR body, treat the comments and imperative instructions inside @.github/pull_request_template.md as rules to follow, not just layout. Agent harnesses may strip HTML comments from copies of that file injected into context, so read .github/pull_request_template.md from disk before writing a PR body to make sure you see every comment rule diff --git a/CONTRIBUTING.md b/CONTRIBUTING.md index 0202965ec4b..d995ddcc87e 100644 --- a/CONTRIBUTING.md +++ b/CONTRIBUTING.md @@ -322,7 +322,7 @@ npm run build ## Submitting Your PR 1. **Push your branch**: `git push origin your-feature-branch` -2. **Create a PR**: Go to GitHub and open a pull request against the current daily OSS branch, named `litellm_oss_daily_YYYY_MM_DD`. A fresh one is cut each weekday, so pick the most recent from the [branch list](https://github.com/BerriAI/litellm/branches/all?query=litellm_oss_daily). Do not target `main`. +2. **Create a PR**: Go to GitHub and open a pull request against [`litellm_internal_staging`](https://github.com/BerriAI/litellm/tree/litellm_internal_staging), which is the default base branch. Do not target `main`. 3. **Fill out the PR template**: Provide clear description of changes 4. **Wait for review**: Maintainers will review and provide feedback 5. **Address feedback**: Make requested changes and push updates diff --git a/litellm-rust/CLAUDE.md b/litellm-rust/CLAUDE.md index 519b1d205ef..0659e63df39 100644 --- a/litellm-rust/CLAUDE.md +++ b/litellm-rust/CLAUDE.md @@ -62,6 +62,13 @@ Not allowed in `core`: Python owns rollout state and fallback while Rust is being introduced. Rust paths must be off by default until parity tests prove equivalence with Python. +A new provider/route may instead be implemented rust-only with no Python +reference; then the Python interface is a thin dispatch that calls Rust with no +fallback, and you state the rust-only choice explicitly in the PR. Either way +the Python side stays minimal (it only marshals inputs and calls the Rust +interface), never add a per-route feature flag, and never push provider +dispatch into `litellm/main.py`; put it in a thin dispatch class under +`litellm/llms///`. ## Production Bar diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 402d16715a3..ce28f737334 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -13,13 +13,13 @@ dependencies = [ [[package]] name = "async-trait" -version = "0.1.89" +version = "0.1.91" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb" +checksum = "ae36dc4177970ef04fde5178d3e2429882def40e57a451f919c098f72baa6cec" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.0", ] [[package]] @@ -113,7 +113,7 @@ dependencies = [ "bytes-utils", "fastrand", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "percent-encoding", "pin-project-lite", "tracing", @@ -193,7 +193,7 @@ dependencies = [ "futures-core", "futures-util", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "http-body-util", "percent-encoding", "pin-project-lite", @@ -222,7 +222,7 @@ dependencies = [ "hyper-util", "pin-project-lite", "rustls 0.21.12", - "rustls 0.23.41", + "rustls 0.23.42", "rustls-native-certs", "rustls-pki-types", "tokio", @@ -279,7 +279,7 @@ dependencies = [ "http 0.2.12", "http 1.4.2", "http-body 0.4.6", - "http-body 1.0.1", + "http-body 1.1.0", "http-body-util", "pin-project-lite", "pin-utils", @@ -313,7 +313,7 @@ checksum = "221eaa237ddf1ca79b60d1372aad77e47f9c0ea5b3ce5099da8c61d027dc77b3" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -340,7 +340,7 @@ dependencies = [ "http 0.2.12", "http 1.4.2", "http-body 0.4.6", - "http-body 1.0.1", + "http-body 1.1.0", "http-body-util", "itoa", "num-integer", @@ -392,7 +392,7 @@ dependencies = [ "bytes", "futures-util", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "http-body-util", "hyper 1.10.1", "hyper-util", @@ -427,7 +427,7 @@ dependencies = [ "bytes", "futures-util", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "http-body-util", "mime", "pin-project-lite", @@ -456,9 +456,9 @@ dependencies = [ [[package]] name = "bitflags" -version = "2.13.0" +version = "2.13.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" +checksum = "b588b76d00fde79687d7646a9b5bdf3cc0f655e0bbd080335a95d7e96f3587da" [[package]] name = "block-buffer" @@ -492,9 +492,9 @@ checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" [[package]] name = "bytes" -version = "1.12.0" +version = "1.12.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" +checksum = "fc652a48c352aef3ea3aed32080501cf3ef6ed5da78602a020c991775b0aff04" [[package]] name = "bytes-utils" @@ -508,9 +508,9 @@ dependencies = [ [[package]] name = "cc" -version = "1.2.65" +version = "1.3.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e228eec9be7c17ccb640b59b36a5cd805ea2a564a4c5e162c2f659fea30d3b96" +checksum = "c89588d05638b5b4594a3348a2d6c20277e43a7f5c5202b05cc56888475a47b8" dependencies = [ "find-msvc-tools", "jobserver", @@ -526,9 +526,20 @@ checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" [[package]] name = "cfg_aliases" -version = "0.2.1" +version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" +checksum = "f079e83a288787bcd14a6aea84cee5c87a67c5a3e660c30f557a3d24761b3527" + +[[package]] +name = "chacha20" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81" +dependencies = [ + "cfg-if", + "cpufeatures 0.3.0", + "rand_core 0.10.1", +] [[package]] name = "cmake" @@ -655,7 +666,7 @@ checksum = "1ac70aa55017e108007fbaf5aa0f54b021c98f92ff8af59d42eda9da96e3dd4f" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -678,9 +689,9 @@ checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" [[package]] name = "fastrand" -version = "2.4.1" +version = "2.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9f1f227452a390804cdb637b74a86990f2a7d7ba4b7d5693aac9b4dd6defd8d6" +checksum = "da7c62ceae207dd37ea5b845da6a0696c799f85e97da1ab5b7910be3c1c80223" [[package]] name = "find-msvc-tools" @@ -711,9 +722,9 @@ checksum = "42703706b716c37f96a77aea830392ad231f44c9e9a67872fa5548707e11b11c" [[package]] name = "futures-channel" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" +checksum = "262590f4fe6afeb0bc83be1daa64e52657fe185690a958af7f3ad0e92085c5ae" dependencies = [ "futures-core", "futures-sink", @@ -721,44 +732,44 @@ dependencies = [ [[package]] name = "futures-core" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" +checksum = "2cd50c473c80f6d7c3670a752354b8e569b1a7cbfdc0419ec88e5edad85e0dc7" [[package]] name = "futures-io" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" +checksum = "4577ecaa3c4f96589d473f679a71b596316f6641bc350038b962a5daf0085d7a" [[package]] name = "futures-macro" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +checksum = "2d6d3cde68c518367be28956066ddfef33813991b77a55005a69dae04bf3b10b" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] name = "futures-sink" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893" +checksum = "e34418ac499d6305c2fb5ad0ed2f6ac998c5f8ca209b4510f7f94242c647e307" [[package]] name = "futures-task" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393" +checksum = "b231ed28831efb4a61a08580c4bc233ec56bc009f4cd8f52da2c3cb97df0c109" [[package]] name = "futures-util" -version = "0.3.32" +version = "0.3.33" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" +checksum = "a77a90a256fce34da66415271e30f94ee91c57b04b8a2c042d9cf3220179deaa" dependencies = [ "futures-core", "futures-io", @@ -793,20 +804,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "getrandom" -version = "0.3.4" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" -dependencies = [ - "cfg-if", - "js-sys", - "libc", - "r-efi 5.3.0", - "wasip2", - "wasm-bindgen", -] - [[package]] name = "getrandom" version = "0.4.3" @@ -814,8 +811,11 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "300e883d756b2e4ec94e02791f39b04b522276138852cfc41d9fb7e904106099" dependencies = [ "cfg-if", + "js-sys", "libc", - "r-efi 6.0.0", + "r-efi", + "rand_core 0.10.1", + "wasm-bindgen", ] [[package]] @@ -917,9 +917,9 @@ dependencies = [ [[package]] name = "http-body" -version = "1.0.1" +version = "1.1.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1efedce1fb8e6913f23e0c92de8e62cd5b772a67e7b3946df930a62566c93184" +checksum = "ca2a8f2913ee65f60facd6a5905613afaa448497a0230cc41ce022d93290bc2c" dependencies = [ "bytes", "http 1.4.2", @@ -927,14 +927,14 @@ dependencies = [ [[package]] name = "http-body-util" -version = "0.1.3" +version = "0.1.4" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b021d93e26becf5dc7e1b75b1bed1fd93124b374ceb73f43d4d4eafec896a64a" +checksum = "e9f41fd6a08e4d4ec69df65976da761afd5ad5e58a9d4acb46bd1c953a9e3ff2" dependencies = [ "bytes", "futures-core", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "pin-project-lite", ] @@ -995,7 +995,7 @@ dependencies = [ "futures-core", "h2 0.4.15", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "httparse", "httpdate", "itoa", @@ -1029,7 +1029,7 @@ dependencies = [ "http 1.4.2", "hyper 1.10.1", "hyper-util", - "rustls 0.23.41", + "rustls 0.23.42", "rustls-native-certs", "tokio", "tokio-rustls 0.26.4", @@ -1048,13 +1048,13 @@ dependencies = [ "futures-channel", "futures-util", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "hyper 1.10.1", "ipnet", "libc", "percent-encoding", "pin-project-lite", - "socket2 0.6.4", + "socket2 0.6.5", "tokio", "tower-service", "tracing", @@ -1242,12 +1242,12 @@ dependencies = [ "aws-sigv4", "aws-smithy-runtime-api", "aws-types", - "rand 0.8.6", + "rand 0.8.7", "reqwest", "serde", "serde_json", "sha2 0.10.9", - "thiserror 2.0.18", + "thiserror 2.0.19", "tokio", ] @@ -1289,9 +1289,9 @@ checksum = "0e7465ac9959cc2b1404e8e2367b43684a6d13790fe23056cc8c6c5a6b7bcb94" [[package]] name = "memchr" -version = "2.8.2" +version = "2.8.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4" +checksum = "cf8baf1c55e62ffcace7a9f06f4bd9cd3f0c4beb022d3b367256b91b87513d98" [[package]] name = "mime" @@ -1301,9 +1301,9 @@ checksum = "6877bb514081ee2a7ff5ef9de3281f14a4dd4bceac4c09388074a6b5df8a139a" [[package]] name = "mio" -version = "1.2.1" +version = "1.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "02bd0af71c67b473010cbbc60715ee815645a4dc942899111f494b4b737d6fda" +checksum = "30d65c71f1ce40ab09135ce117d742b9f8a19ff91a41a8b57ed50bc2de59c427" dependencies = [ "libc", "wasi", @@ -1378,9 +1378,9 @@ checksum = "19f132c84eca552bf34cab8ec81f1c1dcc229b811638f9d283dceabe58c5569e" [[package]] name = "portable-atomic" -version = "1.13.1" +version = "1.14.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49" +checksum = "3d20d5497ef88037a52ff98267d066e7f11fcc5e99bbfbd58a42336193aacec3" [[package]] name = "potential_utf" @@ -1408,9 +1408,9 @@ dependencies = [ [[package]] name = "proc-macro2" -version = "1.0.106" +version = "1.0.107" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +checksum = "985e7ec9bb745e6ce6535b544d84d6cd6f7ad8bd711c398938ae983b91a766d9" dependencies = [ "unicode-ident", ] @@ -1471,7 +1471,7 @@ dependencies = [ "proc-macro2", "pyo3-macros-backend", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -1483,7 +1483,7 @@ dependencies = [ "heck", "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -1498,9 +1498,9 @@ dependencies = [ "quinn-proto", "quinn-udp", "rustc-hash", - "rustls 0.23.41", - "socket2 0.6.4", - "thiserror 2.0.18", + "rustls 0.23.42", + "socket2 0.6.5", + "thiserror 2.0.19", "tokio", "tracing", "web-time", @@ -1508,20 +1508,21 @@ dependencies = [ [[package]] name = "quinn-proto" -version = "0.11.15" +version = "0.11.16" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4fcb935c5bec503c2f0e306bdd3e58bb9029dcb14fa8d9ac76e3a5256ac0763e" +checksum = "2f4bfc015262b9df63c8845072ce59068853ff5872180c2ce2f13038b970e560" dependencies = [ "bytes", - "getrandom 0.3.4", + "getrandom 0.4.3", "lru-slab", - "rand 0.9.4", + "rand 0.10.2", + "rand_pcg", "ring", "rustc-hash", - "rustls 0.23.41", + "rustls 0.23.42", "rustls-pki-types", "slab", - "thiserror 2.0.18", + "thiserror 2.0.19", "tinyvec", "tracing", "web-time", @@ -1529,33 +1530,27 @@ dependencies = [ [[package]] name = "quinn-udp" -version = "0.5.14" +version = "0.5.15" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "addec6a0dcad8a8d96a771f815f0eaf55f9d1805756410b39f5fa81332574cbd" +checksum = "35a133f956daabe89a61a685c2649f13d82d5aa4bd5d12d1277e1072a21c0694" dependencies = [ "cfg_aliases", "libc", "once_cell", - "socket2 0.6.4", + "socket2 0.6.5", "tracing", - "windows-sys 0.60.2", + "windows-sys 0.61.2", ] [[package]] name = "quote" -version = "1.0.46" +version = "1.0.47" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368" +checksum = "1fbf4db142a473a8d80c26bbf18454ed458bf8d26c8219c331daecfdbd079001" dependencies = [ "proc-macro2", ] -[[package]] -name = "r-efi" -version = "5.3.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" - [[package]] name = "r-efi" version = "6.0.0" @@ -1564,23 +1559,24 @@ checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" [[package]] name = "rand" -version = "0.8.6" +version = "0.8.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a" +checksum = "22f6172bdec972074665ed81ed53b71da00bfc44b65a753cfde883ec4c702a1a" dependencies = [ "libc", - "rand_chacha 0.3.1", + "rand_chacha", "rand_core 0.6.4", ] [[package]] name = "rand" -version = "0.9.4" +version = "0.10.2" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "44c5af06bb1b7d3216d91932aed5265164bf384dc89cd6ba05cf59a35f5f76ea" +checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" dependencies = [ - "rand_chacha 0.9.0", - "rand_core 0.9.5", + "chacha20", + "getrandom 0.4.3", + "rand_core 0.10.1", ] [[package]] @@ -1593,16 +1589,6 @@ dependencies = [ "rand_core 0.6.4", ] -[[package]] -name = "rand_chacha" -version = "0.9.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" -dependencies = [ - "ppv-lite86", - "rand_core 0.9.5", -] - [[package]] name = "rand_core" version = "0.6.4" @@ -1614,11 +1600,17 @@ dependencies = [ [[package]] name = "rand_core" -version = "0.9.5" +version = "0.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + +[[package]] +name = "rand_pcg" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "caa0f4137e1c0a72f4c651489402276c8e8e1cf081f3b0ba156d2cbeef09e86a" dependencies = [ - "getrandom 0.3.4", + "rand_core 0.10.1", ] [[package]] @@ -1640,7 +1632,7 @@ dependencies = [ "futures-util", "h2 0.4.15", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "http-body-util", "hyper 1.10.1", "hyper-rustls 0.27.9", @@ -1650,7 +1642,7 @@ dependencies = [ "percent-encoding", "pin-project-lite", "quinn", - "rustls 0.23.41", + "rustls 0.23.42", "rustls-pki-types", "serde", "serde_json", @@ -1686,9 +1678,9 @@ dependencies = [ [[package]] name = "rustc-hash" -version = "2.1.2" +version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "94300abf3f1ae2e2b8ffb7b58043de3d399c73fa6f4b73826402a5c457614dbe" +checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" [[package]] name = "rustc_version" @@ -1713,9 +1705,9 @@ dependencies = [ [[package]] name = "rustls" -version = "0.23.41" +version = "0.23.42" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "6b92b125634d9b795e7beca796cc790df15a7fb38323bf3196fda83292d06b1f" +checksum = "3c54fcab019b409d04215d3a17cb438fd7fbf192ee61461f20f4fe18704bc138" dependencies = [ "aws-lc-rs", "once_cell", @@ -1740,9 +1732,9 @@ dependencies = [ [[package]] name = "rustls-pki-types" -version = "1.14.1" +version = "1.15.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "30a7197ae7eb376e574fe940d068c30fe0462554a3ddbe4eca7838e049c937a9" +checksum = "764899a24af3980067ee14bc143654f297b22eaebfe3c7b6b211920a5a59b046" dependencies = [ "web-time", "zeroize", @@ -1772,9 +1764,9 @@ dependencies = [ [[package]] name = "rustversion" -version = "1.0.22" +version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" [[package]] name = "ryu" @@ -1832,9 +1824,9 @@ checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" [[package]] name = "serde" -version = "1.0.228" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +checksum = "4148590afebada386688f18773da617792bf2ef03ffc1e4cbd2b1d45b023e0ba" dependencies = [ "serde_core", "serde_derive", @@ -1842,22 +1834,22 @@ dependencies = [ [[package]] name = "serde_core" -version = "1.0.228" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +checksum = "67dca2c9c51e58a4791a4b1ed58308b39c64224d349a935ab5039aa360942a48" dependencies = [ "serde_derive", ] [[package]] name = "serde_derive" -version = "1.0.228" +version = "1.0.229" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +checksum = "e7a5d71263a5a7d47b41f6b3f06ba276f10cc18b0931f1799f710578e2309348" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.0", ] [[package]] @@ -1898,9 +1890,9 @@ dependencies = [ [[package]] name = "sha1" -version = "0.10.6" +version = "0.10.7" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +checksum = "a978451301f4db1d02937a4ab3ccce137717b81826e79b7d49ffe3244a13c3b8" dependencies = [ "cfg-if", "cpufeatures 0.2.17", @@ -1959,9 +1951,9 @@ dependencies = [ [[package]] name = "socket2" -version = "0.6.4" +version = "0.6.5" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51" +checksum = "c3d1e2c7f27f8d4cb10542a02c49005dbd6e93095799d6f3be745fae9f8fedd4" dependencies = [ "libc", "windows-sys 0.61.2", @@ -1981,9 +1973,20 @@ checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" [[package]] name = "syn" -version = "2.0.118" +version = "2.0.119" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1b9ae57f904213ebb649ce6895b8a66c66f0203b9319718f69a5612a065b1422" +checksum = "872831b642d1a07999a962a351ed35b955ea2cfc8f3862091e2a240a84f17297" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "syn" +version = "3.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f2fac314a64dc9a36e61a9eb4261a5e9bbfbc922b27e518af97bc32b926cf967" dependencies = [ "proc-macro2", "quote", @@ -2007,7 +2010,7 @@ checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -2027,11 +2030,11 @@ dependencies = [ [[package]] name = "thiserror" -version = "2.0.18" +version = "2.0.19" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +checksum = "09a43598840e33d5b0331f38c5e30d13bb11c11210a4b58f0d9b18a5a5eefcd9" dependencies = [ - "thiserror-impl 2.0.18", + "thiserror-impl 2.0.19", ] [[package]] @@ -2042,18 +2045,18 @@ checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] name = "thiserror-impl" -version = "2.0.18" +version = "2.0.19" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +checksum = "43cbfe0cf76104d42a574802844187e84a305e531ed54455f11fbde0f10541cd" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 3.0.0", ] [[package]] @@ -2098,9 +2101,9 @@ dependencies = [ [[package]] name = "tinyvec" -version = "1.11.0" +version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3e61e67053d25a4e82c844e8424039d9745781b3fc4f32b8d55ed50f5f667ef3" +checksum = "bb4ebadaa0af04fab11ae01eb5f9fdb5f9c5b875506e210e71c07873528baa7f" dependencies = [ "tinyvec_macros", ] @@ -2113,28 +2116,28 @@ checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" [[package]] name = "tokio" -version = "1.52.3" +version = "1.53.0" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" +checksum = "d988bcd52dbe076d3d46903332f58c912b87a2c49b1428419a5845154762ffee" dependencies = [ "bytes", "libc", "mio", "pin-project-lite", - "socket2 0.6.4", + "socket2 0.6.5", "tokio-macros", "windows-sys 0.61.2", ] [[package]] name = "tokio-macros" -version = "2.7.0" +version = "2.7.1" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" +checksum = "6328af13490e73a9b4694030fafd93f8c8c6a9dede33e821c3fc63eddf8042ba" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -2153,7 +2156,7 @@ version = "0.26.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" dependencies = [ - "rustls 0.23.41", + "rustls 0.23.42", "tokio", ] @@ -2165,7 +2168,7 @@ checksum = "edc5f74e248dc973e0dbb7b74c7e0d6fcc301c694ff50049504004ef4d0cdcd9" dependencies = [ "futures-util", "log", - "rustls 0.23.41", + "rustls 0.23.42", "rustls-native-certs", "rustls-pki-types", "tokio", @@ -2212,7 +2215,7 @@ dependencies = [ "bytes", "futures-util", "http 1.4.2", - "http-body 1.0.1", + "http-body 1.1.0", "pin-project-lite", "tower", "tower-layer", @@ -2252,7 +2255,7 @@ checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -2282,8 +2285,8 @@ dependencies = [ "http 1.4.2", "httparse", "log", - "rand 0.8.6", - "rustls 0.23.41", + "rand 0.8.7", + "rustls 0.23.42", "rustls-pki-types", "sha1", "thiserror 1.0.69", @@ -2375,15 +2378,6 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" -[[package]] -name = "wasip2" -version = "1.0.4+wasi-0.2.12" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487" -dependencies = [ - "wit-bindgen", -] - [[package]] name = "wasm-bindgen" version = "0.2.126" @@ -2426,7 +2420,7 @@ dependencies = [ "bumpalo", "proc-macro2", "quote", - "syn", + "syn 2.0.119", "wasm-bindgen-shared", ] @@ -2474,9 +2468,9 @@ dependencies = [ [[package]] name = "webpki-roots" -version = "1.0.8" +version = "1.0.9" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "bf85cb06032201fa7c6f829d7db5a7e5aa45bcc0655327713065f6f0576731bf" +checksum = "7dcd9d09a39985f5344844e66b0c530a33843579125f23e21e9f0f220850f22a" dependencies = [ "rustls-pki-types", ] @@ -2493,16 +2487,7 @@ version = "0.52.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" dependencies = [ - "windows-targets 0.52.6", -] - -[[package]] -name = "windows-sys" -version = "0.60.2" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "f2f500e4d28234f72040990ec9d39e3a6b950f9f22d3dba18416c35882612bcb" -dependencies = [ - "windows-targets 0.53.5", + "windows-targets", ] [[package]] @@ -2520,31 +2505,14 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" dependencies = [ - "windows_aarch64_gnullvm 0.52.6", - "windows_aarch64_msvc 0.52.6", - "windows_i686_gnu 0.52.6", - "windows_i686_gnullvm 0.52.6", - "windows_i686_msvc 0.52.6", - "windows_x86_64_gnu 0.52.6", - "windows_x86_64_gnullvm 0.52.6", - "windows_x86_64_msvc 0.52.6", -] - -[[package]] -name = "windows-targets" -version = "0.53.5" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "4945f9f551b88e0d65f3db0bc25c33b8acea4d9e41163edf90dcd0b19f9069f3" -dependencies = [ - "windows-link", - "windows_aarch64_gnullvm 0.53.1", - "windows_aarch64_msvc 0.53.1", - "windows_i686_gnu 0.53.1", - "windows_i686_gnullvm 0.53.1", - "windows_i686_msvc 0.53.1", - "windows_x86_64_gnu 0.53.1", - "windows_x86_64_gnullvm 0.53.1", - "windows_x86_64_msvc 0.53.1", + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", ] [[package]] @@ -2553,102 +2521,48 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" -[[package]] -name = "windows_aarch64_gnullvm" -version = "0.53.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "a9d8416fa8b42f5c947f8482c43e7d89e73a173cead56d044f6a56104a6d1b53" - [[package]] name = "windows_aarch64_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" -[[package]] -name = "windows_aarch64_msvc" -version = "0.53.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b9d782e804c2f632e395708e99a94275910eb9100b2114651e04744e9b125006" - [[package]] name = "windows_i686_gnu" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" -[[package]] -name = "windows_i686_gnu" -version = "0.53.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "960e6da069d81e09becb0ca57a65220ddff016ff2d6af6a223cf372a506593a3" - [[package]] name = "windows_i686_gnullvm" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" -[[package]] -name = "windows_i686_gnullvm" -version = "0.53.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "fa7359d10048f68ab8b09fa71c3daccfb0e9b559aed648a8f95469c27057180c" - [[package]] name = "windows_i686_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" -[[package]] -name = "windows_i686_msvc" -version = "0.53.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1e7ac75179f18232fe9c285163565a57ef8d3c89254a30685b57d83a38d326c2" - [[package]] name = "windows_x86_64_gnu" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" -[[package]] -name = "windows_x86_64_gnu" -version = "0.53.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "9c3842cdd74a865a8066ab39c8a7a473c0778a3f29370b5fd6b4b9aa7df4a499" - [[package]] name = "windows_x86_64_gnullvm" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" -[[package]] -name = "windows_x86_64_gnullvm" -version = "0.53.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "0ffa179e2d07eee8ad8f57493436566c7cc30ac536a3379fdf008f47f6bb7ae1" - [[package]] name = "windows_x86_64_msvc" version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" -[[package]] -name = "windows_x86_64_msvc" -version = "0.53.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "d6bbff5f0aada427a1e5a6da5f1f98158182f26556f345ac9e04d36d0ebed650" - -[[package]] -name = "wit-bindgen" -version = "0.57.1" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" - [[package]] name = "writeable" version = "0.6.3" @@ -2680,28 +2594,28 @@ checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", "synstructure", ] [[package]] name = "zerocopy" -version = "0.8.52" +version = "0.8.54" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "ce1022995ff5ff5d841ad7d994facc23098cd40152f2c1d11cd607c6f530653f" +checksum = "b7cbbc0a705a0fd05cc3676525980d2bf5a9bc4adac6d6475209a7887cf59d19" dependencies = [ "zerocopy-derive", ] [[package]] name = "zerocopy-derive" -version = "0.8.52" +version = "0.8.54" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "1ae7f38b72ec2a254e2b87ef277cf2cd4fb97cbebf944faa6f33354da0867930" +checksum = "e2e817b7b52d0c7358d3246da9d69935ebb18116b2b102b4230dac079b4862f5" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] @@ -2721,7 +2635,7 @@ checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", "synstructure", ] @@ -2761,11 +2675,11 @@ checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" dependencies = [ "proc-macro2", "quote", - "syn", + "syn 2.0.119", ] [[package]] name = "zmij" -version = "1.0.21" +version = "1.0.23" source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" +checksum = "29666d0abbfad1e3dc4dcf6144730dd3a3ab225bbbdac83319345b1b44ccfc1b" diff --git a/litellm-rust/Cargo.toml b/litellm-rust/Cargo.toml index a3baa33e6cf..6d63be05d00 100644 --- a/litellm-rust/Cargo.toml +++ b/litellm-rust/Cargo.toml @@ -7,7 +7,8 @@ members = [ resolver = "2" [workspace.package] -edition = "2021" +edition = "2024" +rust-version = "1.88" license = "MIT" repository = "https://github.com/BerriAI/litellm" diff --git a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md index c7980b11147..ed44dc4c729 100644 --- a/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md +++ b/litellm-rust/crates/CODING_STANDARDS/PROVIDER_CODING_STANDARDS.md @@ -39,11 +39,17 @@ Rules for adding or changing an LLM provider/route in `litellm-rust`. OCR (`MIST 19. Every provider transform ships tests for: supported-param filtering, request body shape, response normalization, missing/null fields, bad input, and `*_match_python` fixture parity. 20. Lifecycle/hook tests cover hook order, success + failure callback payloads, pre-call guardrail blocking before any provider I/O, during-call body mutation, and provider-error mapping. -21. Rust paths stay off by default and behind Python parity tests (disabled / enabled-equals-Python / bridge-unavailable fallback) until parity is proven. +21. When a route has a Python reference implementation, the Rust path stays off by default and behind Python parity tests (disabled / enabled-equals-Python / bridge-unavailable fallback) until parity is proven. A new provider/route may instead be implemented rust-only with no Python reference; then the Python interface is a thin dispatch to Rust with no fallback, and tests cover the rust-backed path plus the unavailable-bridge error. State the rust-only choice explicitly in the PR. + +## Python bridge (SDK side) + +22. A Python -> Rust bridge keeps the Python side minimal: the Python interface only marshals inputs and calls the Rust interface, with no transform, handler, or business logic. Aim for well under 100 lines of interface code per route; if the Python grows past that, the logic belongs in Rust. +23. Do not bloat `litellm/main.py`. A route's provider dispatch lives in a thin dispatch class under `litellm/llms///` that calls the Rust bridge; `main.py` only instantiates it and calls its sync/async method. +24. Do not add new feature flags unless explicitly requested. Reuse the existing litellm rust rollout mechanism (`use_litellm_rust`); never introduce a per-route env flag such as `LITELLM_USE_RUST_`. ## Checks before push -22. Run, and keep green: +25. Run, and keep green: ```bash cd litellm-rust cargo fmt --check diff --git a/litellm-rust/crates/ai-gateway/Cargo.toml b/litellm-rust/crates/ai-gateway/Cargo.toml index c15af4cc478..541beabe170 100644 --- a/litellm-rust/crates/ai-gateway/Cargo.toml +++ b/litellm-rust/crates/ai-gateway/Cargo.toml @@ -14,7 +14,7 @@ path = "src/main.rs" required-features = ["server"] [dependencies] -litellm-core.workspace = true +litellm-core = { workspace = true, features = ["bedrock-auth"] } # reqwest (rustls + json) is used by io/ocr and ships realtime logs to the # Python proxy callbacks API. reqwest.workspace = true diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/common_utils.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/common_utils.rs new file mode 100644 index 00000000000..270d5c2d97a --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/common_utils.rs @@ -0,0 +1,48 @@ +use std::collections::BTreeMap; + +use litellm_core::CoreResult; +use litellm_core::audio_transcription::transformation::AudioTranscriptionProviderConfig; +use litellm_core::error::CoreError; +use litellm_core::providers::bedrock::audio_transcription::BEDROCK_AUDIO_TRANSCRIPTION_CONFIG; +use serde_json::{Map, Value}; + +pub(super) fn audio_transcription_provider_config( + provider: &str, +) -> Option<&'static dyn AudioTranscriptionProviderConfig> { + match provider { + "bedrock" => Some(&BEDROCK_AUDIO_TRANSCRIPTION_CONFIG), + _ => None, + } +} + +pub(super) fn string_headers( + headers: Option>, +) -> CoreResult> { + headers + .unwrap_or_default() + .into_iter() + .map(|(key, value)| { + value + .as_str() + .map(|value| (key.clone(), value.to_string())) + .ok_or_else(|| { + CoreError::InvalidRequest(format!( + "audio transcription extra_headers.{key} must be a string" + )) + }) + }) + .collect() +} + +pub(super) fn has_header(headers: &BTreeMap, name: &str) -> bool { + headers.keys().any(|key| key.eq_ignore_ascii_case(name)) +} + +pub(super) fn truncate_error_body(body: &str) -> String { + let truncated: String = body.chars().take(256).collect(); + if truncated.chars().count() == body.chars().count() { + truncated + } else { + format!("{truncated}... (truncated)") + } +} diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/handler.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/handler.rs new file mode 100644 index 00000000000..33c13550f58 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/handler.rs @@ -0,0 +1,89 @@ +use std::time::SystemTime; + +use litellm_core::CoreResult; +use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth; +use litellm_core::error::CoreError; +use litellm_core::providers::bedrock::audio_transcription::aws_auth_config; +use litellm_core::providers::bedrock::aws_base::{resolve_credentials, sign_bedrock_post}; +use serde_json::Value; + +use super::common_utils::truncate_error_body; +use super::types::ProviderAudioTranscriptionRequest; +use crate::client::http_client; + +pub(crate) async fn execute_audio_transcription_provider_call( + request: ProviderAudioTranscriptionRequest, +) -> CoreResult { + let body = serde_json::to_vec(&request.body).map_err(|error| { + CoreError::InvalidRequest(format!("invalid audio request body: {error}")) + })?; + let mut request_builder = http_client().post(&request.url).body(body.clone()); + 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 = request_builder + .send() + .await + .map_err(|error| CoreError::Network(error.to_string()))?; + let status = response.status(); + let text = response + .text() + .await + .map_err(|error| CoreError::Network(error.to_string()))?; + if !status.is_success() { + return Err(CoreError::Http { + status: status.as_u16(), + body: truncate_error_body(&text), + }); + } + let response_json: Value = serde_json::from_str(&text).map_err(|error| { + CoreError::InvalidResponse(format!("invalid audio response JSON: {error}")) + })?; + Ok(request + .config + .transform_transcription_response(&request.model, response_json)? + .into_json()) +} + +pub(crate) async fn sign_request( + request: &ProviderAudioTranscriptionRequest, + optional_params: &serde_json::Map, +) -> CoreResult { + let env_lookup = environment_lookup; + let auth = request + .config + .auth_strategy(&request.model, optional_params, &env_lookup)?; + let body = serde_json::to_vec(&request.body).map_err(|error| { + CoreError::InvalidRequest(format!("invalid audio request body: {error}")) + })?; + let mut headers = super::common_utils::string_headers(None)?; + headers.insert("Content-Type".to_string(), "application/json".to_string()); + headers.extend(request.upstream_headers.iter().cloned()); + match auth { + AudioTranscriptionAuth::Bearer => {} + AudioTranscriptionAuth::AwsSigV4 { region, .. } => { + let credentials = + resolve_credentials(aws_auth_config(optional_params, &env_lookup), &env_lookup) + .await?; + headers.extend(sign_bedrock_post( + &request.url, + &body, + &headers, + ®ion, + &credentials, + SystemTime::now(), + )?); + } + } + Ok(ProviderAudioTranscriptionRequest { + upstream_headers: headers.into_iter().collect(), + ..request.clone() + }) +} + +pub(super) fn environment_lookup(key: &str) -> Option { + std::env::var(key).ok() +} diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs new file mode 100644 index 00000000000..8b6896f3846 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/hooks.rs @@ -0,0 +1,300 @@ +use std::future::Future; +use std::pin::Pin; + +use litellm_core::CoreResult; +use litellm_core::audio_transcription::transformation::AudioTranscriptionAuth; +use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; +use litellm_core::error::CoreError; +use serde_json::{Map, Value, json}; + +use super::common_utils::{audio_transcription_provider_config, has_header, string_headers}; +use super::handler::sign_request; +use super::types::{PreparedAudioTranscriptionRequest, ProviderAudioTranscriptionRequest}; +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 AudioTranscriptionLifecycleHooks { + logger_runner: CustomLoggerRunner, + guardrail_runner: CustomGuardrailRunner, + request_metadata: RequestMetadata, +} + +type AudioFuture<'a, T> = Pin> + Send + 'a>>; +type AudioLogFuture<'a> = Pin + Send + 'a>>; + +impl AudioTranscriptionLifecycleHooks { + 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: PreparedAudioTranscriptionRequest, + ) -> CoreResult { + if self.guardrail_runner.is_empty() { + return Ok(request); + } + let (guardrail_request, _) = self + .guardrail_runner + .run_pre_call( + &guardrail_context(&self.request_metadata), + GuardrailRequest::new(json!({ + "model": request.model, + "custom_llm_provider": request.custom_llm_provider, + "audio": request.audio, + "optional_params": request.optional_params, + })), + ) + .await + .map_err(guardrail_error_to_core_error)?; + let Value::Object(mut data) = guardrail_request.data else { + return Err(CoreError::InvalidRequest( + "audio transcription pre_call guardrail must return an object".to_string(), + )); + }; + let audio = data.remove("audio").ok_or_else(|| { + CoreError::InvalidRequest("audio transcription guardrail removed audio".to_string()) + })?; + let optional_params = match data.remove("optional_params") { + Some(Value::Object(value)) => value, + Some(_) => { + return Err(CoreError::InvalidRequest( + "audio transcription optional_params must be an object".to_string(), + )); + } + None => Map::new(), + }; + Ok(PreparedAudioTranscriptionRequest { + audio, + optional_params, + ..request + }) + } + + async fn prepare_provider_request( + &self, + request: PreparedAudioTranscriptionRequest, + ) -> CoreResult { + let config = audio_transcription_provider_config(&request.custom_llm_provider) + .ok_or_else(|| CoreError::InvalidProvider(request.custom_llm_provider.clone()))?; + let env_lookup = super::handler::environment_lookup; + let headers = string_headers(request.extra_headers)?; + let url = config.complete_url( + request.api_base.as_deref(), + &request.model, + &request.optional_params, + &env_lookup, + )?; + let filtered_params = config.map_transcription_params(&request.optional_params); + let body = config.transform_transcription_request( + &request.model, + request.audio, + filtered_params, + )?; + let auth = config.auth_strategy(&request.model, &request.optional_params, &env_lookup)?; + let mut upstream_headers = headers.into_iter().collect::>(); + if matches!(auth, AudioTranscriptionAuth::Bearer) + && !has_header( + &upstream_headers + .iter() + .cloned() + .collect::>(), + "authorization", + ) + && let Some(api_key) = request.api_key.as_deref() + { + upstream_headers.push(("Authorization".to_string(), format!("Bearer {api_key}"))); + } + let provider_request = ProviderAudioTranscriptionRequest { + model: request.model, + config, + url, + body: body.body, + upstream_headers, + timeout: request.timeout, + }; + let provider_request = self.run_during_call_guardrails(provider_request).await?; + sign_request(&provider_request, &request.optional_params).await + } + + async fn run_during_call_guardrails( + &self, + request: ProviderAudioTranscriptionRequest, + ) -> CoreResult { + if self.guardrail_runner.is_empty() { + return Ok(request); + } + let (guardrail_request, _) = self + .guardrail_runner + .run_during_call( + &guardrail_context(&self.request_metadata), + GuardrailRequest::new(json!({ + "model": request.model, + "custom_llm_provider": "bedrock", + "url": request.url, + "body": request.body, + })), + ) + .await + .map_err(guardrail_error_to_core_error)?; + let Value::Object(mut data) = guardrail_request.data else { + return Err(CoreError::InvalidRequest( + "audio transcription during_call guardrail must return an object".to_string(), + )); + }; + let body = data.remove("body").ok_or_else(|| { + CoreError::InvalidRequest("audio transcription guardrail removed body".to_string()) + })?; + Ok(ProviderAudioTranscriptionRequest { body, ..request }) + } + + fn 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, + } + } +} + +impl CallLifecycleHooks + for AudioTranscriptionLifecycleHooks +{ + type PreCallFuture<'a> = AudioFuture<'a, PreparedAudioTranscriptionRequest>; + type DuringCallFuture<'a> = AudioFuture<'a, ProviderAudioTranscriptionRequest>; + type SuccessFuture<'a> = AudioLogFuture<'a>; + type FailureFuture<'a> = AudioLogFuture<'a>; + + fn async_pre_call_hook<'a>( + &'a self, + _context: &'a CallLifecycleContext, + request: PreparedAudioTranscriptionRequest, + ) -> 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: PreparedAudioTranscriptionRequest, + ) -> Self::DuringCallFuture<'a> { + Box::pin(async move { self.prepare_provider_request(request).await }) + } + + 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; + } + self.logger_runner + .async_log_success_event( + &ModelCallDetails::from_standard_logging_payload( + self.logging_payload(context, timing), + ), + &CallbackValue::new("audio_transcription", response.clone()), + CallbackTiming::new(timing.start_time, timing.end_time), + ) + .await; + }) + } + + fn async_log_failure_event<'a>( + &'a self, + context: &'a CallLifecycleContext, + error: &'a CoreError, + 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(), + }; + self.logger_runner + .async_log_failure_event( + &ModelCallDetails::from_standard_logging_payload( + self.logging_payload(context, timing), + ) + .with_failure_error(logging_error.clone()), + Some(&CallbackValue::new( + "error", + json!({"message": logging_error.message, "kind": logging_error.kind}), + )), + CallbackTiming::new(timing.start_time, timing.end_time), + ) + .await; + }) + } +} + +fn guardrail_context(metadata: &RequestMetadata) -> GuardrailContext { + GuardrailContext { + call_type: CallType::Other("audio_transcription".to_string()), + 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 guardrail_error_to_core_error(error: GuardrailError) -> CoreError { + CoreError::InvalidRequest(format!("{}: {}", error.kind, error.message)) +} + +fn core_error_kind(error: &CoreError) -> &'static str { + match error { + CoreError::Auth(_) => "AuthError", + CoreError::InvalidProvider(_) => "InvalidProvider", + CoreError::InvalidRequest(_) => "InvalidRequest", + CoreError::InvalidType { .. } => "InvalidType", + CoreError::MissingField(_) => "MissingField", + CoreError::Http { .. } => "HttpError", + CoreError::InvalidResponse(_) => "InvalidResponse", + CoreError::Network(_) => "NetworkError", + CoreError::Routing(_) => "RoutingError", + } +} diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs new file mode 100644 index 00000000000..5d33d912c40 --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/mod.rs @@ -0,0 +1,25 @@ +use litellm_core::CoreResult; +use litellm_core::call_lifecycle::CallLifecycle; +use serde_json::Value; + +mod common_utils; +mod handler; +mod hooks; +mod prepare; +mod types; + +pub use types::AudioTranscriptionRequest; + +use handler::execute_audio_transcription_provider_call; +use prepare::{PreparedAudioTranscriptionCall, prepare_audio_transcription_call}; + +pub async fn audio_transcription(request: AudioTranscriptionRequest<'_>) -> CoreResult { + let PreparedAudioTranscriptionCall { request, hooks } = + prepare_audio_transcription_call(request); + CallLifecycle::default() + .run_request(request, &hooks, execute_audio_transcription_provider_call) + .await +} + +#[cfg(test)] +mod tests; diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs new file mode 100644 index 00000000000..a475d58635f --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/prepare.rs @@ -0,0 +1,55 @@ +use std::sync::atomic::{AtomicU64, Ordering}; +use std::time::{SystemTime, UNIX_EPOCH}; + +use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; + +use super::hooks::AudioTranscriptionLifecycleHooks; +use super::types::{AudioTranscriptionRequest, PreparedAudioTranscriptionRequest}; +use crate::integrations::custom_guardrail::CustomGuardrailRunner; +use crate::integrations::custom_logger::CustomLoggerRunner; + +pub(crate) struct PreparedAudioTranscriptionCall { + pub(crate) request: PreparedAudioTranscriptionRequest, + pub(crate) hooks: AudioTranscriptionLifecycleHooks, +} + +pub(crate) fn prepare_audio_transcription_call( + request: AudioTranscriptionRequest<'_>, +) -> PreparedAudioTranscriptionCall { + let call_id = request + .litellm_call_id + .map(str::to_string) + .unwrap_or_else(new_audio_transcription_call_id); + let provider_info = get_custom_llm_provider(request.model, request.custom_llm_provider) + .unwrap_or(CustomLlmProvider { + model: request.model, + custom_llm_provider: "bedrock", + }); + PreparedAudioTranscriptionCall { + request: PreparedAudioTranscriptionRequest { + model: provider_info.model.to_string(), + custom_llm_provider: provider_info.custom_llm_provider.to_string(), + litellm_call_id: call_id, + audio: request.audio, + 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: request.optional_params, + timeout: request.timeout, + }, + hooks: AudioTranscriptionLifecycleHooks::new( + CustomLoggerRunner::new(request.callbacks), + CustomGuardrailRunner::new(request.guardrails), + request.request_metadata, + ), + } +} + +fn new_audio_transcription_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_or(0, |duration| duration.as_nanos()); + format!("audio-transcription-{timestamp}-{sequence}") +} diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs new file mode 100644 index 00000000000..5df04708b7d --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/tests.rs @@ -0,0 +1,53 @@ +use std::io::{Read, Write}; +use std::net::TcpListener; +use std::thread; + +use serde_json::{Map, json}; + +use super::{AudioTranscriptionRequest, audio_transcription}; + +#[tokio::test] +async fn bedrock_request_is_signed_and_contains_audio() { + let listener = TcpListener::bind("127.0.0.1:0").expect("listener"); + let address = listener.local_addr().expect("address"); + let server = thread::spawn(move || { + let (mut stream, _) = listener.accept().expect("connection"); + let mut request = Vec::new(); + let mut buffer = [0_u8; 16_384]; + let count = stream.read(&mut buffer).expect("request"); + request.extend_from_slice(&buffer[..count]); + let request = String::from_utf8_lossy(&request); + assert!(request.contains("POST /model/mistral.voxtral-mini-3b-2507/converse")); + assert!(request.contains("authorization: AWS4-HMAC-SHA256")); + assert!(request.contains("x-amz-date:")); + assert!(request.contains("\"bytes\":\"AQI=\"")); + assert!(request.contains("Transcribe the audio. Respond with only the transcript.")); + let response = b"HTTP/1.1 200 OK\r\nContent-Type: application/json\r\nContent-Length: 53\r\nConnection: close\r\n\r\n{\"output\":{\"message\":{\"content\":[{\"text\":\"hello\"}]}}}"; + stream.write_all(response).expect("response"); + }); + + let optional_params = Map::from_iter([ + ("aws_access_key_id".to_string(), json!("access-key")), + ("aws_secret_access_key".to_string(), json!("secret-key")), + ("aws_region_name".to_string(), json!("us-east-1")), + ]); + let api_base = format!("http://{address}"); + let response = audio_transcription(AudioTranscriptionRequest { + model: "mistral.voxtral-mini-3b-2507", + audio: json!({"data": "AQI=", "format": "wav", "filename": "audio.wav"}), + api_key: None, + api_base: Some(&api_base), + custom_llm_provider: Some("bedrock"), + extra_headers: None, + optional_params, + timeout: None, + callbacks: Vec::new(), + guardrails: Vec::new(), + request_metadata: Default::default(), + litellm_call_id: None, + }) + .await + .expect("transcription"); + assert_eq!(response, json!({"text": "hello"})); + server.join().expect("server"); +} diff --git a/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs new file mode 100644 index 00000000000..9697aa98b0a --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/audio_transcription/types.rs @@ -0,0 +1,58 @@ +use std::sync::Arc; +use std::time::Duration; + +use litellm_core::audio_transcription::transformation::AudioTranscriptionProviderConfig; +use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleRequest}; +use serde_json::{Map, Value}; + +use crate::integrations::custom_guardrail::CustomGuardrail; +use crate::integrations::custom_logger::CustomLogger; +use crate::integrations::types::RequestMetadata; + +pub struct AudioTranscriptionRequest<'a> { + pub model: &'a str, + pub audio: Value, + pub api_key: Option<&'a str>, + pub api_base: Option<&'a str>, + pub custom_llm_provider: Option<&'a str>, + pub extra_headers: Option>, + pub optional_params: Map, + pub timeout: Option, + pub callbacks: Vec>, + pub guardrails: Vec>, + pub request_metadata: RequestMetadata, + pub litellm_call_id: Option<&'a str>, +} + +pub(crate) struct PreparedAudioTranscriptionRequest { + pub(crate) model: String, + pub(crate) custom_llm_provider: String, + pub(crate) litellm_call_id: String, + pub(crate) audio: 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 PreparedAudioTranscriptionRequest { + fn lifecycle_context(&self) -> CallLifecycleContext { + CallLifecycleContext::new( + "audio_transcription", + self.model.clone(), + self.custom_llm_provider.clone(), + self.litellm_call_id.clone(), + ) + } +} + +#[derive(Clone)] +pub(crate) struct ProviderAudioTranscriptionRequest { + pub(crate) model: String, + pub(crate) config: &'static dyn AudioTranscriptionProviderConfig, + pub(crate) url: String, + pub(crate) body: Value, + pub(crate) upstream_headers: Vec<(String, String)>, + pub(crate) timeout: Option, +} diff --git a/litellm-rust/crates/ai-gateway/src/auth/mod.rs b/litellm-rust/crates/ai-gateway/src/auth/mod.rs index 438a0513057..b09d8285c3a 100644 --- a/litellm-rust/crates/ai-gateway/src/auth/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/auth/mod.rs @@ -9,9 +9,9 @@ //! runs during extraction, before the handler body. Routes never re-implement it. use axum::extract::FromRequestParts; +use axum::http::StatusCode; use axum::http::header::AUTHORIZATION; use axum::http::request::Parts; -use axum::http::StatusCode; use sha2::{Digest, Sha256}; use subtle::ConstantTimeEq; diff --git a/litellm-rust/crates/ai-gateway/src/ocr/client.rs b/litellm-rust/crates/ai-gateway/src/client.rs similarity index 60% rename from litellm-rust/crates/ai-gateway/src/ocr/client.rs rename to litellm-rust/crates/ai-gateway/src/client.rs index 79cc7816227..ff2606f0229 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/client.rs +++ b/litellm-rust/crates/ai-gateway/src/client.rs @@ -1,13 +1,13 @@ use std::sync::OnceLock; use std::time::Duration; -const OCR_TIMEOUT_SECS: u64 = 600; +const HTTP_CLIENT_TIMEOUT_SECS: u64 = 600; -pub(super) fn http_client() -> &'static reqwest::Client { +pub(crate) fn http_client() -> &'static reqwest::Client { static CLIENT: OnceLock = OnceLock::new(); CLIENT.get_or_init(|| { reqwest::Client::builder() - .timeout(Duration::from_secs(OCR_TIMEOUT_SECS)) + .timeout(Duration::from_secs(HTTP_CLIENT_TIMEOUT_SECS)) .build() .expect("failed to build reqwest client") }) diff --git a/litellm-rust/crates/ai-gateway/src/io/audio_transcription.rs b/litellm-rust/crates/ai-gateway/src/io/audio_transcription.rs new file mode 100644 index 00000000000..80d9e401a5f --- /dev/null +++ b/litellm-rust/crates/ai-gateway/src/io/audio_transcription.rs @@ -0,0 +1 @@ +pub use crate::audio_transcription::{AudioTranscriptionRequest, audio_transcription}; diff --git a/litellm-rust/crates/ai-gateway/src/io/messages.rs b/litellm-rust/crates/ai-gateway/src/io/messages.rs index b784d2b62a1..86170e45678 100644 --- a/litellm-rust/crates/ai-gateway/src/io/messages.rs +++ b/litellm-rust/crates/ai-gateway/src/io/messages.rs @@ -1 +1 @@ -pub use crate::messages::{messages, MessagesRequest}; +pub use crate::messages::{MessagesRequest, messages}; diff --git a/litellm-rust/crates/ai-gateway/src/io/mod.rs b/litellm-rust/crates/ai-gateway/src/io/mod.rs index 9cbfa568121..6129a808965 100644 --- a/litellm-rust/crates/ai-gateway/src/io/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/io/mod.rs @@ -1,3 +1,4 @@ +pub mod audio_transcription; pub mod messages; pub mod ocr; pub mod realtime; diff --git a/litellm-rust/crates/ai-gateway/src/io/ocr.rs b/litellm-rust/crates/ai-gateway/src/io/ocr.rs index 55e02839c4e..2fc82f0b61f 100644 --- a/litellm-rust/crates/ai-gateway/src/io/ocr.rs +++ b/litellm-rust/crates/ai-gateway/src/io/ocr.rs @@ -1 +1 @@ -pub use crate::ocr::{ocr, OcrRequest}; +pub use crate::ocr::{OcrRequest, ocr}; diff --git a/litellm-rust/crates/ai-gateway/src/io/realtime.rs b/litellm-rust/crates/ai-gateway/src/io/realtime.rs index 40a38c1579a..845e7bf9527 100644 --- a/litellm-rust/crates/ai-gateway/src/io/realtime.rs +++ b/litellm-rust/crates/ai-gateway/src/io/realtime.rs @@ -15,16 +15,16 @@ use std::time::Duration; use futures_util::stream::{SplitSink, SplitStream}; use futures_util::{Sink, SinkExt, Stream, StreamExt}; +use litellm_core::CoreResult; use litellm_core::error::CoreError; use litellm_core::realtime::transformation::RealtimeProviderConfig; use litellm_core::realtime::types::RealtimeEvent; -use litellm_core::CoreResult; use tokio::net::TcpStream; -use tokio_tungstenite::tungstenite::client::IntoClientRequest; -use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION; -use tokio_tungstenite::tungstenite::http::HeaderValue; use tokio_tungstenite::tungstenite::Message; -use tokio_tungstenite::{connect_async, MaybeTlsStream, WebSocketStream}; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::HeaderValue; +use tokio_tungstenite::tungstenite::http::header::AUTHORIZATION; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async}; use litellm_core::providers::openai::realtime::transformation::OPENAI_REALTIME_CONFIG; @@ -113,7 +113,7 @@ pub(crate) async fn read_event(upstream_rx: &mut UpstreamRx) -> CoreResult { return Err(CoreError::Network( "upstream closed before first event".to_string(), - )) + )); } _ => continue, } diff --git a/litellm-rust/crates/ai-gateway/src/io/realtime_pool.rs b/litellm-rust/crates/ai-gateway/src/io/realtime_pool.rs index bf8041f31d7..4a1a3cd1166 100644 --- a/litellm-rust/crates/ai-gateway/src/io/realtime_pool.rs +++ b/litellm-rust/crates/ai-gateway/src/io/realtime_pool.rs @@ -28,11 +28,11 @@ use std::sync::{Arc, Mutex}; use std::time::{Duration, Instant}; use futures_util::StreamExt; -use litellm_core::realtime::types::RealtimeEvent; use litellm_core::CoreResult; +use litellm_core::realtime::types::RealtimeEvent; use crate::io::realtime::{ - dial_upstream, read_event, resolve_api_key, UpstreamRx, UpstreamTx, UpstreamWs, + UpstreamRx, UpstreamTx, UpstreamWs, dial_upstream, read_event, resolve_api_key, }; /// Default target warm sockets per key when pooling is enabled. @@ -473,8 +473,8 @@ pub fn upstream_key( /// unhealthy — we'd rather discard and fresh-dial than hand over a socket in an /// unexpected state. `Pending` (the healthy case) returns `false`. fn is_dead(rx: &mut UpstreamRx) -> bool { - use futures_util::task::noop_waker_ref; use futures_util::Stream; + use futures_util::task::noop_waker_ref; use std::pin::Pin; use std::task::{Context, Poll}; @@ -523,15 +523,15 @@ mod tests { )) .await; while let Some(Ok(msg)) = ws.next().await { - if let Message::Text(text) = msg { - if text.contains("response.create") { - for frame in [ - r#"{"type":"response.created"}"#, - r#"{"type":"response.output_audio.delta","delta":"AAAA"}"#, - r#"{"type":"response.done"}"#, - ] { - let _ = ws.send(Message::Text(frame.to_string())).await; - } + if let Message::Text(text) = msg + && text.contains("response.create") + { + for frame in [ + r#"{"type":"response.created"}"#, + r#"{"type":"response.output_audio.delta","delta":"AAAA"}"#, + r#"{"type":"response.done"}"#, + ] { + let _ = ws.send(Message::Text(frame.to_string())).await; } } } 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 ae6ad150bcf..9b51019f4bc 100644 --- a/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs +++ b/litellm-rust/crates/ai-gateway/src/io/responses_ws.rs @@ -10,19 +10,18 @@ use litellm_core::responses::websocket::ResponsesWebSocketProviderConfig; use litellm_core::{CoreError, CoreResult}; use tokio::net::TcpStream; use tokio::sync::Mutex; -use tokio_tungstenite::tungstenite::client::IntoClientRequest; -use tokio_tungstenite::tungstenite::http::header::{HeaderName, AUTHORIZATION}; -use tokio_tungstenite::tungstenite::http::HeaderValue; use tokio_tungstenite::tungstenite::Message; -use tokio_tungstenite::{connect_async, MaybeTlsStream, WebSocketStream}; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::HeaderValue; +use tokio_tungstenite::tungstenite::http::header::{AUTHORIZATION, HeaderName}; +use tokio_tungstenite::{MaybeTlsStream, WebSocketStream, connect_async}; use crate::constants::{ DEFAULT_RESPONSES_WS_CONNECT_TIMEOUT_SECS, DEFAULT_RESPONSES_WS_IDLE_TIMEOUT_SECS, }; 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"; +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; @@ -83,8 +82,8 @@ impl ResponsesWebSocketConnection { } pub async fn recv_text(&self) -> CoreResult> { - let mut socket = self.socket.lock().await; - let Some(socket) = socket.as_mut() else { + let mut socket_guard = self.socket.lock().await; + let Some(socket) = socket_guard.as_mut() else { return Ok(None); }; match socket.next().await { @@ -456,9 +455,11 @@ mod tests { assert_eq!(fourth.event_type, ResponsesWsEventType::ResponseCompleted); let observed: Vec<_> = observed_rx.collect().await; assert_eq!(observed.len(), 4); - assert!(observed - .iter() - .all(|event| event.event_type != ResponsesWsEventType::ResponseCreate)); + assert!( + observed + .iter() + .all(|event| event.event_type != ResponsesWsEventType::ResponseCreate) + ); } #[tokio::test] diff --git a/litellm-rust/crates/ai-gateway/src/lib.rs b/litellm-rust/crates/ai-gateway/src/lib.rs index 25aac3c495b..c44d661c29e 100644 --- a/litellm-rust/crates/ai-gateway/src/lib.rs +++ b/litellm-rust/crates/ai-gateway/src/lib.rs @@ -11,6 +11,8 @@ //! binary turns on. The `python-config` feature additionally pulls in [`python`] //! for the load-time config reader. +pub mod audio_transcription; +mod client; pub mod io; pub mod messages; pub mod ocr; diff --git a/litellm-rust/crates/ai-gateway/src/main.rs b/litellm-rust/crates/ai-gateway/src/main.rs index f9ce97801d3..da3a486d4ee 100644 --- a/litellm-rust/crates/ai-gateway/src/main.rs +++ b/litellm-rust/crates/ai-gateway/src/main.rs @@ -11,7 +11,7 @@ use std::sync::Arc; -use litellm_ai_gateway::io::realtime_pool::{upstream_key, PoolConfig, RealtimePool}; +use litellm_ai_gateway::io::realtime_pool::{PoolConfig, RealtimePool, upstream_key}; use litellm_ai_gateway::routes; use litellm_ai_gateway::state::AppState; use litellm_core::router::{Deployment, LiteLLMParams, Router}; diff --git a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs b/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs index 33894d0ee64..4b906155665 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/common_utils.rs @@ -1,8 +1,8 @@ -use litellm_core::error::{json_type_name, CoreError}; +use litellm_core::CoreResult; +use litellm_core::error::{CoreError, json_type_name}; use litellm_core::messages::transformation::AnthropicMessagesProviderConfig; use litellm_core::providers::anthropic::messages::transformation::ANTHROPIC_MESSAGES_CONFIG; use litellm_core::providers::azure_ai::messages::transformation::AZURE_ANTHROPIC_MESSAGES_CONFIG; -use litellm_core::CoreResult; use serde_json::{Map, Value}; use crate::constants::MESSAGES_ERROR_BODY_MAX_CHARS; diff --git a/litellm-rust/crates/ai-gateway/src/messages/handler.rs b/litellm-rust/crates/ai-gateway/src/messages/handler.rs index d3b9d3b3fba..90c12367f50 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/handler.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/handler.rs @@ -1,5 +1,5 @@ -use litellm_core::error::CoreError; use litellm_core::CoreResult; +use litellm_core::error::CoreError; use serde_json::Value; use super::client::http_client; diff --git a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs index 6176f9cb67f..624c3598fb0 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/prepare.rs @@ -1,7 +1,7 @@ -use litellm_core::messages::transformation::MessagesAuthStrategy; -use litellm_core::routing_utils::provider::{get_custom_llm_provider, CustomLlmProvider}; use litellm_core::CoreError; use litellm_core::CoreResult; +use litellm_core::messages::transformation::MessagesAuthStrategy; +use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; use super::common_utils::{has_header, messages_provider_config, string_headers}; use super::types::{MessagesRequest, ProviderMessagesRequest}; diff --git a/litellm-rust/crates/ai-gateway/src/messages/tests.rs b/litellm-rust/crates/ai-gateway/src/messages/tests.rs index 9b1cc45aacb..a2d0f6fae23 100644 --- a/litellm-rust/crates/ai-gateway/src/messages/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/messages/tests.rs @@ -1,14 +1,14 @@ use std::time::Duration; use litellm_core::error::CoreError; -use serde_json::{json, Map, Value}; +use serde_json::{Map, Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; use super::common_utils::{ has_header, messages_provider_config, string_headers, truncate_error_body, }; -use super::{messages, MessagesRequest}; +use super::{MessagesRequest, messages}; async fn read_http_request(socket: &mut TcpStream) -> String { let mut request = Vec::new(); diff --git a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs b/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs index d4b4d9338e7..9bc2818b6e7 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/common_utils.rs @@ -1,11 +1,11 @@ use std::net::IpAddr; use std::time::{Duration, Instant}; -use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; use base64::Engine; +use base64::engine::general_purpose::STANDARD as BASE64_STANDARD; +use litellm_core::CoreResult; use litellm_core::error::CoreError; use litellm_core::ocr::transformation::OcrProviderConfig; -use litellm_core::CoreResult; use reqwest::Url; use serde_json::{Map, Value}; @@ -18,7 +18,7 @@ use litellm_core::providers::vertex_ai::ocr::transformation::{ VERTEX_AI_DEEPSEEK_OCR_CONFIG, VERTEX_AI_OCR_CONFIG, }; -use super::client::http_client; +use crate::client::http_client; const ERROR_BODY_MAX_CHARS: usize = 256; const AZURE_DOCUMENT_INTELLIGENCE_POLL_TIMEOUT_SECS: u64 = 120; diff --git a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs index 4d93c2a25db..1de34eb400e 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/handler.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/handler.rs @@ -1,11 +1,11 @@ +use litellm_core::CoreResult; use litellm_core::error::CoreError; use litellm_core::ocr::transformation::OcrResponseHandling; -use litellm_core::CoreResult; use serde_json::Value; -use super::client::http_client; use super::common_utils::{poll_document_intelligence, truncate_error_body}; use super::types::ProviderOcrRequest; +use crate::client::http_client; pub(crate) async fn execute_ocr_provider_call(request: ProviderOcrRequest) -> CoreResult { let mut request_builder = http_client().post(&request.url).json(&request.body); diff --git a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs index 6be74ed2714..ffe2e0122c0 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/hooks.rs @@ -1,11 +1,11 @@ use std::future::Future; use std::pin::Pin; +use litellm_core::CoreResult; use litellm_core::call_lifecycle::{CallLifecycleContext, CallLifecycleHooks, CallLifecycleTiming}; use litellm_core::error::CoreError; use litellm_core::ocr::transformation::OcrAuthStrategy; -use litellm_core::CoreResult; -use serde_json::{json, Map, Value}; +use serde_json::{Map, Value, json}; use super::common_utils::{ convert_document_url_to_data_uri, has_header, ocr_provider_config, string_headers, @@ -292,7 +292,7 @@ fn parse_ocr_pre_call_guardrail_request( Some(_) => { return Err(CoreError::InvalidRequest( "OCR pre_call guardrail optional_params must be an object".to_string(), - )) + )); } None => Map::new(), }; diff --git a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs index b54ee39b21d..c4c13e2300c 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/mod.rs @@ -1,8 +1,7 @@ -use litellm_core::call_lifecycle::CallLifecycle; use litellm_core::CoreResult; +use litellm_core::call_lifecycle::CallLifecycle; use serde_json::Value; -mod client; mod common_utils; mod handler; mod hooks; @@ -12,7 +11,7 @@ mod types; pub use types::OcrRequest; use handler::execute_ocr_provider_call; -use prepare::{prepare_ocr_call, PreparedOcrCall}; +use prepare::{PreparedOcrCall, prepare_ocr_call}; pub async fn ocr(request: OcrRequest<'_>) -> CoreResult { let PreparedOcrCall { request, hooks } = prepare_ocr_call(request); diff --git a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs index 5a4b350a4c4..6231393c889 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/prepare.rs @@ -1,7 +1,7 @@ use std::sync::atomic::{AtomicU64, Ordering}; use std::time::{SystemTime, UNIX_EPOCH}; -use litellm_core::routing_utils::provider::{get_custom_llm_provider, CustomLlmProvider}; +use litellm_core::routing_utils::provider::{CustomLlmProvider, get_custom_llm_provider}; use super::hooks::OcrLifecycleHooks; use super::types::{OcrRequest, PreparedOcrRequest}; diff --git a/litellm-rust/crates/ai-gateway/src/ocr/tests.rs b/litellm-rust/crates/ai-gateway/src/ocr/tests.rs index 35747dc6985..bb2a6b06501 100644 --- a/litellm-rust/crates/ai-gateway/src/ocr/tests.rs +++ b/litellm-rust/crates/ai-gateway/src/ocr/tests.rs @@ -3,12 +3,12 @@ use std::time::Duration; use litellm_core::error::CoreError; use litellm_core::ocr::transformation::OcrResponseHandling; -use serde_json::{json, Map, Value}; +use serde_json::{Map, Value, json}; use tokio::io::{AsyncReadExt, AsyncWriteExt}; use tokio::net::{TcpListener, TcpStream}; use super::common_utils::{has_header, ocr_provider_config, string_headers, truncate_error_body}; -use super::{ocr, OcrRequest}; +use super::{OcrRequest, ocr}; use crate::integrations::custom_guardrail::{ CustomGuardrail, GuardrailContext, GuardrailDecision, GuardrailError, GuardrailEventHook, GuardrailFuture, GuardrailRequest, @@ -228,19 +228,23 @@ fn truncate_error_body_does_not_split_multibyte_chars() { #[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!( + 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("vertex_ai", "deepseek-ocr-maas") + .expect("vertex deepseek config resolves") + .supported_ocr_params() + .contains(&"temperature") + ); assert!(ocr_provider_config("openai", "gpt-4o").is_none()); } diff --git a/litellm-rust/crates/ai-gateway/src/python/config.rs b/litellm-rust/crates/ai-gateway/src/python/config.rs index 54b7a53bafa..c028d3d6b51 100644 --- a/litellm-rust/crates/ai-gateway/src/python/config.rs +++ b/litellm-rust/crates/ai-gateway/src/python/config.rs @@ -7,9 +7,9 @@ //! //! Compiled only under the `python-config` feature. +use litellm_core::CoreResult; use litellm_core::error::CoreError; use litellm_core::router::{Deployment, Router}; -use litellm_core::CoreResult; use pyo3::prelude::*; use crate::gil; diff --git a/litellm-rust/crates/ai-gateway/src/routes/health.rs b/litellm-rust/crates/ai-gateway/src/routes/health.rs index 15c67fea325..c64ca3a7199 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/health.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/health.rs @@ -1,8 +1,8 @@ //! Health probes. Simple-route template: a `router()` plus its handlers, in one file. +use axum::Router; use axum::http::StatusCode; use axum::routing::get; -use axum::Router; use crate::state::AppState; 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 933386282fa..a34b2edd7b8 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/mod.rs @@ -2,13 +2,13 @@ mod service; +use axum::Router; use axum::body::Body; use axum::extract::{Json, State}; -use axum::http::header::{HeaderMap, HeaderValue, CACHE_CONTROL, CONTENT_TYPE}; use axum::http::StatusCode; +use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE, HeaderMap, HeaderValue}; use axum::response::{IntoResponse, Response}; use axum::routing::post; -use axum::Router; use litellm_core::CoreError; use serde_json::{Map, Value}; @@ -125,9 +125,9 @@ mod tests { use std::sync::Arc; use axum::body::Body; - use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE}; use axum::http::Request; use axum::http::StatusCode; + use axum::http::header::{CACHE_CONTROL, CONTENT_TYPE}; use litellm_core::router::{Deployment, LiteLLMParams, Router as ModelRouter}; use serde_json::json; use tokio::io::{AsyncReadExt, AsyncWriteExt}; @@ -439,8 +439,8 @@ mod tests { .await .expect("response body reads"); assert_eq!( - serde_json::from_slice::(&response_body).expect("error is json") - ["error"]["message"], + serde_json::from_slice::(&response_body).expect("error is json")["error"] + ["message"], "messages provider request failed" ); server.await.expect("upstream task completes"); diff --git a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs index 7f00123ca39..75ed26e5be8 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/messages/service.rs @@ -5,7 +5,7 @@ use litellm_core::{CoreError, CoreResult}; use serde_json::{Map, Value}; use crate::constants::ANTHROPIC_MESSAGES_PROVIDER; -use crate::messages::{execute_messages, MessagesRequest}; +use crate::messages::{MessagesRequest, execute_messages}; pub(crate) enum MessagesResponse { Json(Value), diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs index c3f929f5f0b..f9144ad1fdb 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/mod.rs @@ -6,17 +6,17 @@ mod service; -use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; use std::time::{SystemTime, UNIX_EPOCH}; use crate::io::realtime_pool::RealtimePool; +use axum::Router; use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade}; use axum::extract::{Query, State}; use axum::http::StatusCode; use axum::response::Response; use axum::routing::get; -use axum::Router; use futures_util::{SinkExt, StreamExt}; use litellm_core::realtime::types::RealtimeEvent; use litellm_core::router::Router as ModelRouter; diff --git a/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs b/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs index d6c31edd454..4ae8cfe7379 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/realtime/service.rs @@ -9,12 +9,12 @@ use std::time::Duration; -use crate::io::realtime_pool::{upstream_key, RealtimePool}; +use crate::io::realtime_pool::{RealtimePool, upstream_key}; use futures_util::{Sink, Stream}; +use litellm_core::CoreResult; use litellm_core::error::CoreError; use litellm_core::realtime::types::RealtimeEvent; use litellm_core::router::Router; -use litellm_core::CoreResult; /// Select a deployment for `model` and splice the client stream to the provider. /// diff --git a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs index bdaffc97afb..a94853e106d 100644 --- a/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs +++ b/litellm-rust/crates/ai-gateway/src/routes/responses/mod.rs @@ -1,15 +1,15 @@ mod service; -use std::sync::atomic::{AtomicU64, Ordering}; use std::sync::Arc; +use std::sync::atomic::{AtomicU64, Ordering}; use std::time::{SystemTime, UNIX_EPOCH}; +use axum::Router; use axum::extract::ws::{Message, WebSocket, WebSocketUpgrade}; use axum::extract::{Query, State}; use axum::http::StatusCode; use axum::response::Response; use axum::routing::get; -use axum::Router; use futures_util::{Sink, SinkExt, StreamExt}; use litellm_core::responses::types::{ResponsesErrorFrame, ResponsesWsEvent, ResponsesWsEventType}; use litellm_core::router::Router as ModelRouter; diff --git a/litellm-rust/crates/core/src/audio_transcription/mod.rs b/litellm-rust/crates/core/src/audio_transcription/mod.rs new file mode 100644 index 00000000000..ec2fbb969a6 --- /dev/null +++ b/litellm-rust/crates/core/src/audio_transcription/mod.rs @@ -0,0 +1,2 @@ +pub mod transformation; +pub mod types; diff --git a/litellm-rust/crates/core/src/audio_transcription/transformation.rs b/litellm-rust/crates/core/src/audio_transcription/transformation.rs new file mode 100644 index 00000000000..eab34c13843 --- /dev/null +++ b/litellm-rust/crates/core/src/audio_transcription/transformation.rs @@ -0,0 +1,57 @@ +use serde_json::{Map, Value}; + +use crate::CoreResult; + +use super::types::{AudioTranscriptionRequestData, AudioTranscriptionResponseData}; + +#[derive(Clone, Debug, PartialEq, Eq)] +pub enum AudioTranscriptionAuth { + Bearer, + AwsSigV4 { + region: String, + service: &'static str, + }, +} + +pub trait AudioTranscriptionProviderConfig: Sync { + fn supported_transcription_params(&self) -> &'static [&'static str]; + + fn map_transcription_params(&self, params: &Map) -> Map { + params + .iter() + .filter(|(key, _)| { + self.supported_transcription_params() + .contains(&key.as_str()) + }) + .map(|(key, value)| (key.clone(), value.clone())) + .collect() + } + + fn transform_transcription_request( + &self, + model: &str, + audio: Value, + optional_params: Map, + ) -> CoreResult; + + fn transform_transcription_response( + &self, + model: &str, + response_json: Value, + ) -> CoreResult; + + fn complete_url( + &self, + api_base: Option<&str>, + model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult; + + fn auth_strategy( + &self, + model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult; +} diff --git a/litellm-rust/crates/core/src/audio_transcription/types.rs b/litellm-rust/crates/core/src/audio_transcription/types.rs new file mode 100644 index 00000000000..3a9e1ecd88c --- /dev/null +++ b/litellm-rust/crates/core/src/audio_transcription/types.rs @@ -0,0 +1,20 @@ +use serde::{Deserialize, Serialize}; +use serde_json::Value; + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct AudioTranscriptionRequestData { + pub body: Value, +} + +#[derive(Clone, Debug, PartialEq, Serialize, Deserialize)] +pub struct AudioTranscriptionResponseData { + pub text: String, +} + +impl AudioTranscriptionResponseData { + pub fn into_json(self) -> Value { + serde_json::json!({ + "text": self.text, + }) + } +} diff --git a/litellm-rust/crates/core/src/caching/in_memory_cache.rs b/litellm-rust/crates/core/src/caching/in_memory_cache.rs index 0ceeedb8b71..45d4bd69b79 100644 --- a/litellm-rust/crates/core/src/caching/in_memory_cache.rs +++ b/litellm-rust/crates/core/src/caching/in_memory_cache.rs @@ -134,8 +134,8 @@ impl InMemoryCache { #[cfg(test)] mod tests { use std::sync::{ - atomic::{AtomicU64, Ordering}, Arc, + atomic::{AtomicU64, Ordering}, }; use super::InMemoryCache; diff --git a/litellm-rust/crates/core/src/lib.rs b/litellm-rust/crates/core/src/lib.rs index 3989fb441bc..51ea19750ea 100644 --- a/litellm-rust/crates/core/src/lib.rs +++ b/litellm-rust/crates/core/src/lib.rs @@ -1,3 +1,4 @@ +pub mod audio_transcription; pub mod caching; pub mod call_lifecycle; pub mod constants; 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 13e79b087c7..6935bb4604b 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 @@ -5,7 +5,7 @@ use crate::messages::types::{ MessageContent, SystemPrompt, }; use crate::providers::anthropic::messages::transformation::{ - non_empty, AnthropicMessagesConfig, ANTHROPIC_MESSAGES_CONFIG, + ANTHROPIC_MESSAGES_CONFIG, AnthropicMessagesConfig, non_empty, }; use serde_json::{Map, Value}; 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 index 060073acd47..eabd15677cc 100644 --- a/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/azure_ai/ocr/transformation.rs @@ -1,9 +1,9 @@ use std::collections::BTreeSet; -use crate::error::{json_type_name, CoreError, CoreResult}; +use crate::error::{CoreError, CoreResult, json_type_name}; use crate::ocr::transformation::{OcrAuthStrategy, OcrProviderConfig, OcrResponseHandling}; use crate::ocr::types::{OcrRequestData, OcrResponseData}; -use serde_json::{json, Map, Value}; +use serde_json::{Map, Value, json}; use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; @@ -206,11 +206,11 @@ pub fn complete_document_intelligence_url( AZURE_DOCUMENT_INTELLIGENCE_API_VERSION ); - if let Some(pages) = optional_params.get("pages") { - if let Some(normalized) = normalize_pages_param(pages)? { - url.push_str("&pages="); - url.push_str(&normalized); - } + if let Some(pages) = optional_params.get("pages") + && let Some(normalized) = normalize_pages_param(pages)? + { + url.push_str("&pages="); + url.push_str(&normalized); } Ok(url) @@ -231,7 +231,7 @@ fn document_url_from_mistral_document(document: &Value) -> CoreResult<&str> { other => { return Err(CoreError::InvalidRequest(format!( "Invalid document type: {other}. Must be 'document_url' or 'image_url'" - ))) + ))); } }; object diff --git a/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs b/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs new file mode 100644 index 00000000000..86eb589e2c0 --- /dev/null +++ b/litellm-rust/crates/core/src/providers/bedrock/audio_transcription.rs @@ -0,0 +1,310 @@ +use serde_json::{Map, Value, json}; + +use crate::audio_transcription::transformation::{ + AudioTranscriptionAuth, AudioTranscriptionProviderConfig, +}; +use crate::audio_transcription::types::{ + AudioTranscriptionRequestData, AudioTranscriptionResponseData, +}; +use crate::error::{CoreError, CoreResult, json_type_name}; + +use super::aws_base::AwsAuthConfig; +use super::constants::{ + AWS_REGION, AWS_REGION_NAME, BEDROCK_RUNTIME_ENDPOINT_TEMPLATE, BEDROCK_SERVICE, + DEFAULT_BEDROCK_REGION, +}; + +const SUPPORTED_PARAMS: &[&str] = &["language", "prompt", "temperature", "response_format"]; + +pub static BEDROCK_AUDIO_TRANSCRIPTION_CONFIG: BedrockAudioTranscriptionConfig = + BedrockAudioTranscriptionConfig; + +pub struct BedrockAudioTranscriptionConfig; + +pub fn bedrock_model_id_and_region(model: &str) -> (String, Option) { + let mut stripped = model; + for prefix in ["bedrock/converse/", "bedrock/", "converse/"] { + if let Some(value) = stripped.strip_prefix(prefix) { + stripped = value; + break; + } + } + let mut region = None; + if let Some((candidate, remainder)) = stripped.split_once('/') + && is_bedrock_region(candidate) + { + region = Some(candidate.to_string()); + stripped = remainder; + } + for prefix in ["nova-2/", "nova/"] { + if let Some(value) = stripped.strip_prefix(prefix) { + stripped = value; + break; + } + } + if region.is_none() { + region = stripped + .strip_prefix("arn:") + .and_then(|value| value.split(':').nth(3)) + .filter(|value| !value.is_empty()) + .map(str::to_string); + } + (stripped.to_string(), region) +} + +fn is_bedrock_region(value: &str) -> bool { + value.len() > 3 + && value.contains('-') + && value + .chars() + .all(|char| char.is_ascii_alphanumeric() || char == '-') +} + +pub fn resolve_bedrock_region( + model_region: Option<&str>, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, +) -> String { + if let Some(region) = optional_params + .get("aws_region_name") + .and_then(Value::as_str) + { + return region.to_string(); + } + if let Some(region) = model_region { + return region.to_string(); + } + env_lookup(AWS_REGION_NAME) + .or_else(|| env_lookup(AWS_REGION)) + .unwrap_or_else(|| DEFAULT_BEDROCK_REGION.to_string()) +} + +fn audio_fields(audio: Value) -> CoreResult<(String, String)> { + let object = audio.as_object().ok_or_else(|| CoreError::InvalidType { + expected: "object", + actual: json_type_name(&audio), + })?; + let data = object + .get("data") + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) + .ok_or(CoreError::MissingField("audio.data"))?; + let format = object + .get("format") + .and_then(Value::as_str) + .filter(|value| matches!(*value, "wav" | "mp3" | "flac" | "ogg")) + .ok_or_else(|| { + CoreError::InvalidRequest("audio.format must be wav, mp3, flac, or ogg".to_string()) + })?; + Ok((data.to_string(), format.to_string())) +} + +fn optional_string<'a>(params: &'a Map, key: &str) -> Option<&'a str> { + params + .get(key) + .and_then(Value::as_str) + .filter(|value| !value.is_empty()) +} + +impl AudioTranscriptionProviderConfig for BedrockAudioTranscriptionConfig { + fn supported_transcription_params(&self) -> &'static [&'static str] { + SUPPORTED_PARAMS + } + + fn transform_transcription_request( + &self, + _model: &str, + audio: Value, + optional_params: Map, + ) -> CoreResult { + let (data, format) = audio_fields(audio)?; + let mut instruction = "Transcribe the audio. Respond with only the transcript.".to_string(); + if let Some(language) = optional_string(&optional_params, "language") { + instruction.push_str(&format!(" The audio language is {language}.")); + } + if let Some(prompt) = optional_string(&optional_params, "prompt") { + instruction.push_str(&format!(" Additional context: {prompt}")); + } + let mut inference_config = Map::from_iter([("maxTokens".to_string(), json!(4096))]); + if let Some(temperature) = optional_params.get("temperature") { + inference_config.insert("temperature".to_string(), temperature.clone()); + } + Ok(AudioTranscriptionRequestData { + body: json!({ + "messages": [{ + "role": "user", + "content": [ + {"audio": {"format": format, "source": {"bytes": data}}}, + {"text": instruction} + ] + }], + "system": [{"text": "You are a transcription assistant."}], + "inferenceConfig": inference_config, + }), + }) + } + + fn transform_transcription_response( + &self, + _model: &str, + response_json: Value, + ) -> CoreResult { + let content = response_json + .get("output") + .and_then(|value| value.get("message")) + .and_then(|value| value.get("content")) + .and_then(Value::as_array) + .ok_or_else(|| { + CoreError::InvalidResponse("Bedrock response has no output content".to_string()) + })?; + let mut text = String::new(); + for block in content { + if let Some(value) = block.get("text").and_then(Value::as_str) { + text.push_str(value); + } + } + Ok(AudioTranscriptionResponseData { text }) + } + + fn complete_url( + &self, + api_base: Option<&str>, + model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + let (model_id, model_region) = bedrock_model_id_and_region(model); + let region = resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup); + let endpoint = optional_params + .get("aws_bedrock_runtime_endpoint") + .and_then(Value::as_str) + .or(api_base) + .map(str::trim) + .filter(|value| !value.is_empty()) + .map(str::to_string) + .unwrap_or_else(|| BEDROCK_RUNTIME_ENDPOINT_TEMPLATE.replace("{region}", ®ion)); + Ok(format!( + "{}/model/{model_id}/converse", + endpoint.trim_end_matches('/') + )) + } + + fn auth_strategy( + &self, + model: &str, + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, + ) -> CoreResult { + let (_, model_region) = bedrock_model_id_and_region(model); + Ok(AudioTranscriptionAuth::AwsSigV4 { + region: resolve_bedrock_region(model_region.as_deref(), optional_params, env_lookup), + service: BEDROCK_SERVICE, + }) + } +} + +pub fn aws_auth_config( + optional_params: &Map, + env_lookup: &dyn Fn(&str) -> Option, +) -> AwsAuthConfig { + let value = |key: &str| { + optional_params + .get(key) + .and_then(Value::as_str) + .map(str::to_string) + }; + let env = |key: &str| env_lookup(key); + AwsAuthConfig { + access_key_id: value("aws_access_key_id").or_else(|| env("AWS_ACCESS_KEY_ID")), + secret_access_key: value("aws_secret_access_key").or_else(|| env("AWS_SECRET_ACCESS_KEY")), + session_token: value("aws_session_token").or_else(|| env("AWS_SESSION_TOKEN")), + region_name: value("aws_region_name").or_else(|| env(AWS_REGION_NAME)), + session_name: value("aws_session_name").or_else(|| env("AWS_SESSION_NAME")), + profile_name: value("aws_profile_name").or_else(|| env("AWS_PROFILE_NAME")), + role_name: value("aws_role_name").or_else(|| env("AWS_ROLE_NAME")), + web_identity_token: value("aws_web_identity_token") + .or_else(|| env("AWS_WEB_IDENTITY_TOKEN")), + sts_endpoint: value("aws_sts_endpoint").or_else(|| env("AWS_STS_ENDPOINT")), + external_id: value("aws_external_id").or_else(|| env("AWS_EXTERNAL_ID")), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + fn no_env(_: &str) -> Option { + None + } + + #[test] + fn request_matches_python_shape() { + let params = Map::from_iter([ + ("language".to_string(), json!("en")), + ("prompt".to_string(), json!("Speaker names")), + ("temperature".to_string(), json!(0)), + ("timestamp_granularities".to_string(), json!(["word"])), + ]); + let params = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG.map_transcription_params(¶ms); + let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG + .transform_transcription_request( + "mistral.voxtral-mini-3b-2507", + json!({"data": "AQI=", "format": "wav", "filename": "sample.wav"}), + params, + ) + .expect("request"); + assert_eq!( + result.body, + json!({ + "messages": [{ + "role": "user", + "content": [ + {"audio": {"format": "wav", "source": {"bytes": "AQI="}}}, + {"text": "Transcribe the audio. Respond with only the transcript. The audio language is en. Additional context: Speaker names"} + ] + }], + "system": [{"text": "You are a transcription assistant."}], + "inferenceConfig": {"maxTokens": 4096, "temperature": 0} + }) + ); + } + + #[test] + fn response_concatenates_content_blocks() { + let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG + .transform_transcription_response( + "model", + json!({"output": {"message": {"content": [{"text": "hello "}, {"text": "world"}]}}}), + ) + .expect("response"); + assert_eq!(result.text, "hello world"); + assert_eq!(result.into_json(), json!({"text": "hello world"})); + } + + #[test] + fn invalid_audio_is_rejected() { + let result = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG.transform_transcription_request( + "model", + json!({"data": "AQI="}), + Map::new(), + ); + assert!(result.is_err()); + } + + #[test] + fn region_and_url_precedence_match_python() { + let params = Map::from_iter([("aws_region_name".to_string(), json!("eu-west-1"))]); + let url = BEDROCK_AUDIO_TRANSCRIPTION_CONFIG + .complete_url( + None, + "bedrock/us-east-1/mistral.voxtral-mini-3b-2507", + ¶ms, + &no_env, + ) + .expect("url"); + assert_eq!( + url, + "https://bedrock-runtime.eu-west-1.amazonaws.com/model/mistral.voxtral-mini-3b-2507/converse" + ); + } +} diff --git a/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs b/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs index 82d1e8fdf91..dc036a3cf21 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/aws_base.rs @@ -5,10 +5,10 @@ use std::time::{SystemTime, UNIX_EPOCH}; use crate::caching::in_memory_cache::InMemoryCache; use crate::error::{CoreError, CoreResult}; -use aws_credential_types::provider::ProvideCredentials; use aws_credential_types::Credentials; +use aws_credential_types::provider::ProvideCredentials; use aws_sigv4::http_request::{ - sign, SignableBody, SignableRequest, SigningParams, SigningSettings, + SignableBody, SignableRequest, SigningParams, SigningSettings, sign, }; use aws_sigv4::sign::v4; use aws_smithy_runtime_api::client::identity::Identity; @@ -52,7 +52,7 @@ pub struct AwsAuthConfig { } impl AwsAuthConfig { - fn with_environment(self, env_lookup: &dyn Fn(&str) -> Option) -> Self { + fn with_environment(self, env_lookup: &(dyn Fn(&str) -> Option + Sync)) -> Self { Self { access_key_id: self.access_key_id.or_else(|| env_lookup(AWS_ACCESS_KEY_ID)), secret_access_key: self @@ -144,7 +144,7 @@ fn same_role_arns(target: &str, caller: &str) -> bool { pub fn classify_auth( config: AwsAuthConfig, - env_lookup: &dyn Fn(&str) -> Option, + env_lookup: &(dyn Fn(&str) -> Option + Sync), ) -> AwsAuthFlow { let config = config.with_environment(env_lookup); if let (Some(token), Some(role), Some(session_name)) = ( @@ -194,7 +194,7 @@ pub fn classify_auth( pub async fn resolve_credentials( config: AwsAuthConfig, - env_lookup: &dyn Fn(&str) -> Option, + env_lookup: &(dyn Fn(&str) -> Option + Sync), ) -> CoreResult { let resolved = config.clone().with_environment(env_lookup); let flow = classify_auth(config, env_lookup); @@ -368,11 +368,11 @@ async fn is_already_running_as_role(role: &str, config: &AwsAuthConfig) -> CoreR if let (Ok(current_role), Ok(token_file)) = ( std::env::var(AWS_ROLE_ARN), std::env::var(AWS_WEB_IDENTITY_TOKEN_FILE), - ) { - if !token_file.is_empty() { - return Ok(same_role_arns(role, ¤t_role)); - } + ) && !token_file.is_empty() + { + return Ok(same_role_arns(role, ¤t_role)); } + let mut loader = aws_config::defaults(aws_config::BehaviorVersion::latest()); if let Some(region) = config.region_name.clone() { loader = loader.region(aws_types::region::Region::new(region)); @@ -639,7 +639,9 @@ mod tests { ); assert_eq!( signed.get("Authorization").map(String::as_str), - Some("AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20240102/us-east-1/bedrock/aws4_request, SignedHeaders=content-type;host;x-amz-date;x-amz-security-token, Signature=55c027ef47527d3ad63f1735f9d099efdbc99f296ff914bd94e727e24ec0e464") + Some( + "AWS4-HMAC-SHA256 Credential=AKIDEXAMPLE/20240102/us-east-1/bedrock/aws4_request, SignedHeaders=content-type;host;x-amz-date;x-amz-security-token, Signature=55c027ef47527d3ad63f1735f9d099efdbc99f296ff914bd94e727e24ec0e464" + ) ); } diff --git a/litellm-rust/crates/core/src/providers/bedrock/constants.rs b/litellm-rust/crates/core/src/providers/bedrock/constants.rs index a08ae9de146..785295207e7 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/constants.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/constants.rs @@ -2,6 +2,7 @@ pub const AWS_ACCESS_KEY_ID: &str = "AWS_ACCESS_KEY_ID"; pub const AWS_SECRET_ACCESS_KEY: &str = "AWS_SECRET_ACCESS_KEY"; pub const AWS_SESSION_TOKEN: &str = "AWS_SESSION_TOKEN"; pub const AWS_REGION_NAME: &str = "AWS_REGION_NAME"; +pub const AWS_REGION: &str = "AWS_REGION"; pub const AWS_SESSION_NAME: &str = "AWS_SESSION_NAME"; pub const AWS_PROFILE_NAME: &str = "AWS_PROFILE_NAME"; pub const AWS_ROLE_NAME: &str = "AWS_ROLE_NAME"; @@ -12,3 +13,6 @@ pub const AWS_STS_ENDPOINT: &str = "AWS_STS_ENDPOINT"; pub const AWS_EXTERNAL_ID: &str = "AWS_EXTERNAL_ID"; pub const BEDROCK_SERVICE: &str = "bedrock"; pub const DEFAULT_SESSION_NAME_PREFIX: &str = "litellm-session"; +pub const DEFAULT_BEDROCK_REGION: &str = "us-west-2"; +pub const BEDROCK_RUNTIME_ENDPOINT_TEMPLATE: &str = + "https://bedrock-runtime.{region}.amazonaws.com"; diff --git a/litellm-rust/crates/core/src/providers/bedrock/mod.rs b/litellm-rust/crates/core/src/providers/bedrock/mod.rs index 8027260a78b..b09675ad7dd 100644 --- a/litellm-rust/crates/core/src/providers/bedrock/mod.rs +++ b/litellm-rust/crates/core/src/providers/bedrock/mod.rs @@ -2,5 +2,7 @@ //! with Python's `BaseAWSLLM`; the broader core purity guidance is reconciled //! separately. +#[cfg(feature = "bedrock-auth")] +pub mod audio_transcription; pub mod aws_base; mod constants; diff --git a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs b/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs index 1a33bc1e951..dc720cc4244 100644 --- a/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/mistral/ocr/transformation.rs @@ -1,4 +1,4 @@ -use crate::error::{json_type_name, CoreError, CoreResult}; +use crate::error::{CoreError, CoreResult, json_type_name}; use crate::ocr::transformation::OcrProviderConfig; use crate::ocr::types::{OcrRequestData, OcrResponseData}; use serde_json::{Map, Value}; diff --git a/litellm-rust/crates/core/src/providers/openai/realtime/transformation.rs b/litellm-rust/crates/core/src/providers/openai/realtime/transformation.rs index 626e4014ff9..b3f6b03b28a 100644 --- a/litellm-rust/crates/core/src/providers/openai/realtime/transformation.rs +++ b/litellm-rust/crates/core/src/providers/openai/realtime/transformation.rs @@ -1,6 +1,6 @@ +use crate::CoreResult; use crate::realtime::transformation::RealtimeProviderConfig; use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult}; -use crate::CoreResult; /// Default OpenAI API base, used when the caller does not override `api_base`. pub const OPENAI_REALTIME_DEFAULT_API_BASE: &str = "https://api.openai.com"; diff --git a/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs b/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs index ece10971806..e15197c468c 100644 --- a/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs +++ b/litellm-rust/crates/core/src/providers/openai/responses/transformation.rs @@ -1,6 +1,6 @@ -use crate::responses::types::{ResponsesWsEvent, ResponsesWsTransformResult}; -use crate::responses::websocket::{enforce_model, ResponsesWebSocketProviderConfig}; use crate::CoreResult; +use crate::responses::types::{ResponsesWsEvent, ResponsesWsTransformResult}; +use crate::responses::websocket::{ResponsesWebSocketProviderConfig, enforce_model}; pub struct OpenAIResponsesWsConfig; 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 index 8639926c435..6300149c237 100644 --- a/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs +++ b/litellm-rust/crates/core/src/providers/vertex_ai/ocr/transformation.rs @@ -1,7 +1,7 @@ -use crate::error::{json_type_name, CoreError, CoreResult}; +use crate::error::{CoreError, CoreResult, json_type_name}; use crate::ocr::transformation::OcrProviderConfig; use crate::ocr::types::{OcrRequestData, OcrResponseData}; -use serde_json::{json, Map, Value}; +use serde_json::{Map, Value, json}; use crate::providers::mistral::ocr::transformation::MISTRAL_OCR_CONFIG; @@ -140,7 +140,7 @@ fn document_content_item(document: &Value) -> CoreResult { other => { return Err(CoreError::InvalidRequest(format!( "Unsupported document type: {other}. Expected 'image_url' or 'document_url'" - ))) + ))); } }; let url = object diff --git a/litellm-rust/crates/core/src/realtime/transformation.rs b/litellm-rust/crates/core/src/realtime/transformation.rs index a4baa27a6c2..69b88687000 100644 --- a/litellm-rust/crates/core/src/realtime/transformation.rs +++ b/litellm-rust/crates/core/src/realtime/transformation.rs @@ -1,5 +1,5 @@ -use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult}; use crate::CoreResult; +use crate::realtime::types::{RealtimeEvent, RealtimeTransformResult}; pub trait RealtimeProviderConfig { /// Build the upstream WebSocket URL (e.g. `wss://api.openai.com/v1/realtime?model=…`). diff --git a/litellm-rust/crates/core/src/responses/websocket.rs b/litellm-rust/crates/core/src/responses/websocket.rs index 1edffd44985..92dc19627a0 100644 --- a/litellm-rust/crates/core/src/responses/websocket.rs +++ b/litellm-rust/crates/core/src/responses/websocket.rs @@ -1,6 +1,6 @@ +use crate::CoreResult; use crate::constants::{OPENAI_RESPONSES_DEFAULT_API_BASE, OPENAI_RESPONSES_PATH}; use crate::responses::types::{ResponsesWsEvent, ResponsesWsEventType, ResponsesWsTransformResult}; -use crate::CoreResult; pub trait ResponsesWebSocketProviderConfig: Sync { fn supports_native_websocket(&self) -> bool { diff --git a/litellm-rust/crates/python-bridge/CLAUDE.md b/litellm-rust/crates/python-bridge/CLAUDE.md index efa1a554c9c..e5d021ec25b 100644 --- a/litellm-rust/crates/python-bridge/CLAUDE.md +++ b/litellm-rust/crates/python-bridge/CLAUDE.md @@ -17,7 +17,13 @@ Python-compatible dictionaries. - Provider dispatch belongs in Rust route modules such as `litellm_providers::ocr`, not in this PyO3 crate. - Python owns rollout state and fallback. Rust should return errors; Python - decides whether to raise or fall back. + decides whether to raise or fall back. For a rust-only provider/route (no + Python reference), the Python side is a thin dispatch that calls Rust and + raises when the bridge is unavailable, with no fallback. +- Keep the Python interface minimal (well under 100 lines per route): it only + marshals inputs and calls Rust. Do not add per-route feature flags, and do + not put provider dispatch in `litellm/main.py`; it lives in a thin dispatch + class under `litellm/llms///`. ## Data Handling diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 83e163c38f1..20a9ba789ce 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -10,7 +10,7 @@ name = "_native" crate-type = ["cdylib"] [dependencies] -litellm-core.workspace = true +litellm-core = { workspace = true, features = ["bedrock-auth"] } litellm-ai-gateway = { workspace = true, default-features = false } pyo3 = { workspace = true, features = ["extension-module"] } pyo3-async-runtimes.workspace = true diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 1decb789a22..ee9bdd0b81f 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -1,8 +1,11 @@ use std::collections::HashMap; use std::time::Duration; -use litellm_ai_gateway::io::messages::{messages as run_messages, MessagesRequest}; -use litellm_ai_gateway::io::ocr::{ocr as run_ocr, OcrRequest}; +use litellm_ai_gateway::io::audio_transcription::{ + AudioTranscriptionRequest, audio_transcription as run_audio_transcription, +}; +use litellm_ai_gateway::io::messages::{MessagesRequest, messages as run_messages}; +use litellm_ai_gateway::io::ocr::{OcrRequest, ocr as run_ocr}; use litellm_ai_gateway::io::responses_ws::ResponsesWebSocketConnection as RustResponsesWebSocketConnection; use litellm_core::error::CoreError; use pyo3::exceptions::{PyRuntimeError, PyValueError}; @@ -244,6 +247,93 @@ fn aocr( }) } +#[pyfunction] +#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))] +#[allow(clippy::too_many_arguments)] +fn transcription( + py: Python<'_>, + model: String, + audio: Py, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + optional_params: Option>, + timeout_seconds: Option, +) -> PyResult> { + let audio = py_to_json(py, audio.bind(py))?; + let extra_headers = match extra_headers { + Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?), + None => None, + }; + let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; + let timeout = optional_timeout(timeout_seconds); + let result = gil::release_gil(py, || { + pyo3_async_runtimes::tokio::get_runtime().block_on(run_audio_transcription( + AudioTranscriptionRequest { + model: &model, + audio, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + optional_params, + timeout, + callbacks: Vec::new(), + guardrails: Vec::new(), + request_metadata: Default::default(), + litellm_call_id: None, + }, + )) + }); + match result { + Ok(value) => json_to_py(py, value), + Err(err) => Err(core_error_to_pyerr(err)), + } +} + +#[pyfunction] +#[pyo3(signature = (model, audio, api_key=None, api_base=None, custom_llm_provider=None, extra_headers=None, optional_params=None, timeout_seconds=None))] +#[allow(clippy::too_many_arguments)] +fn atranscription( + py: Python<'_>, + model: String, + audio: Py, + api_key: Option, + api_base: Option, + custom_llm_provider: Option, + extra_headers: Option>, + optional_params: Option>, + timeout_seconds: Option, +) -> PyResult> { + let audio = py_to_json(py, audio.bind(py))?; + let extra_headers = match extra_headers { + Some(headers) => Some(optional_object_to_map(py, "extra_headers", Some(headers))?), + None => None, + }; + let optional_params = optional_object_to_map(py, "optional_params", optional_params)?; + let timeout = optional_timeout(timeout_seconds); + pyo3_async_runtimes::tokio::future_into_py(py, async move { + let value = run_audio_transcription(AudioTranscriptionRequest { + model: &model, + audio, + api_key: api_key.as_deref(), + api_base: api_base.as_deref(), + custom_llm_provider: custom_llm_provider.as_deref(), + extra_headers, + optional_params, + timeout, + callbacks: Vec::new(), + guardrails: Vec::new(), + request_metadata: Default::default(), + litellm_call_id: None, + }) + .await + .map_err(core_error_to_pyerr)?; + Python::attach(|py| json_to_py(py, value)) + }) +} + type MarshaledMessagesInputs = (Value, Option>, Option); fn marshal_messages_inputs( @@ -341,6 +431,8 @@ fn gil_stats(py: Python<'_>) -> PyResult> { fn _native(module: &Bound<'_, PyModule>) -> PyResult<()> { module.add_function(wrap_pyfunction!(ocr, module)?)?; module.add_function(wrap_pyfunction!(aocr, module)?)?; + module.add_function(wrap_pyfunction!(transcription, module)?)?; + module.add_function(wrap_pyfunction!(atranscription, module)?)?; module.add_function(wrap_pyfunction!(messages, module)?)?; module.add_function(wrap_pyfunction!(amessages, module)?)?; module.add_class::()?; diff --git a/litellm/caching/disk_cache.py b/litellm/caching/disk_cache.py index d9f65ce949e..af8eb92849f 100644 --- a/litellm/caching/disk_cache.py +++ b/litellm/caching/disk_cache.py @@ -58,12 +58,12 @@ class DiskCache(BaseCache): return return_val def increment_cache(self, key, value: int, **kwargs) -> int: - # get the value - cached_value = self.get_cache(key=key) - init_value = cached_value if isinstance(cached_value, int) else 0 - value = init_value + value - self.set_cache(key, value, **kwargs) - return value + with self.disk_cache.transact(): + cached_value = self.get_cache(key=key) + init_value = cached_value if isinstance(cached_value, int) else 0 + new_value = init_value + value + self.set_cache(key, new_value, **kwargs) + return new_value async def async_get_cache(self, key, **kwargs): return self.get_cache(key=key, **kwargs) @@ -76,12 +76,7 @@ class DiskCache(BaseCache): return return_val async def async_increment(self, key, value: int, **kwargs) -> int: - # get the value - cached_value = await self.async_get_cache(key=key) - init_value = cached_value if isinstance(cached_value, int) else 0 - value = init_value + value - await self.async_set_cache(key, value, **kwargs) - return value + return self.increment_cache(key=key, value=value, **kwargs) def flush_cache(self): self.disk_cache.clear() diff --git a/litellm/caching/in_memory_cache.py b/litellm/caching/in_memory_cache.py index 2ad3f3f11b7..36b477f7a8b 100644 --- a/litellm/caching/in_memory_cache.py +++ b/litellm/caching/in_memory_cache.py @@ -12,6 +12,7 @@ import json import sys import time import heapq +import threading from typing import TYPE_CHECKING, Any, List, Optional if TYPE_CHECKING: @@ -46,6 +47,7 @@ class InMemoryCache(BaseCache): self.cache_dict: dict = {} self.ttl_dict: dict = {} self.expiration_heap: list[tuple[float, str]] = [] + self._increment_lock = threading.Lock() def check_value_size(self, value: Any): """ @@ -223,12 +225,13 @@ class InMemoryCache(BaseCache): return_val.append(val) return return_val - def increment_cache(self, key, value: int, **kwargs) -> int: - # get the value - init_value = self.get_cache(key=key) or 0 - value = init_value + value - self.set_cache(key, value, **kwargs) - return value + def increment_cache(self, key, value: float, **kwargs) -> float: + with self._increment_lock: + # keep read-modify-write atomic + init_value = self.get_cache(key=key) or 0 + value = init_value + value + self.set_cache(key, value, **kwargs) + return value async def async_get_cache(self, key, **kwargs): return self.get_cache(key=key, **kwargs) @@ -241,11 +244,7 @@ class InMemoryCache(BaseCache): return return_val async def async_increment(self, key, value: float, **kwargs) -> float: - # get the value - init_value = await self.async_get_cache(key=key) or 0 - value = init_value + value - await self.async_set_cache(key, value, **kwargs) - return value + return self.increment_cache(key=key, value=value, **kwargs) async def async_increment_pipeline( self, increment_list: List["RedisPipelineIncrementOperation"], **kwargs diff --git a/litellm/experimental_mcp_client/tools.py b/litellm/experimental_mcp_client/tools.py index c65b266bd02..500d226752b 100644 --- a/litellm/experimental_mcp_client/tools.py +++ b/litellm/experimental_mcp_client/tools.py @@ -9,6 +9,7 @@ from openai.types.chat import ChatCompletionToolParam from openai.types.responses.function_tool_param import FunctionToolParam from openai.types.shared_params.function_definition import FunctionDefinition +from litellm.types.llms.anthropic import AnthropicMessagesTool from litellm.types.utils import ChatCompletionMessageToolCall @@ -75,6 +76,20 @@ def transform_mcp_tool_to_openai_responses_api_tool( ) +def transform_mcp_tool_to_anthropic_tool(mcp_tool: MCPTool) -> AnthropicMessagesTool: + """Convert an MCP tool to an Anthropic Messages API tool.""" + from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, + ) + + return AnthropicMessagesTool( + name=mcp_tool.name, + description=mcp_tool.description or "", + input_schema=sanitize_input_schema_for_anthropic(mcp_tool.inputSchema), + type="custom", + ) + + async def load_mcp_tools( session: ClientSession, format: Literal["mcp", "openai"] = "mcp" ) -> Union[List[MCPTool], List[ChatCompletionToolParam]]: diff --git a/litellm/integrations/anthropic_cache_control_hook.py b/litellm/integrations/anthropic_cache_control_hook.py index 94c86e07ff5..faedf8ae1a3 100644 --- a/litellm/integrations/anthropic_cache_control_hook.py +++ b/litellm/integrations/anthropic_cache_control_hook.py @@ -91,7 +91,9 @@ class AnthropicCacheControlHook(CustomPromptManagement): # Pass through non-message injection points for provider-specific handling if remaining_points: - non_default_params["cache_control_injection_points"] = remaining_points + non_default_params["cache_control_injection_points"] = AnthropicCacheControlHook._stamped_as_judged( + remaining_points + ) return model, processed_messages, non_default_params @@ -310,6 +312,35 @@ class AnthropicCacheControlHook(CustomPromptManagement): return ChatCompletionCachedContent(type="ephemeral", ttl=ttl) return ChatCompletionCachedContent(type="ephemeral") + @staticmethod + def _stamped_as_judged(points: list[CacheControlInjectionPoint]) -> list[dict[str, object]]: + """Mark written-back points as having passed the client cache_control judgment. + + Builds copies because config-owned point dicts are shared across + requests; mutating them would leak the stamp into future requests. + """ + return [{**point, "_litellm_judged": True} for point in points] + + @staticmethod + def _should_stand_down( + points: list[CacheControlInjectionPoint], + messages: list[AllMessageValues], + system: str | list | None, + tools: list | None, + ) -> bool: + """Whether configured injection points must yield to client-set cache_control. + + Points that a prior pass over this request already judged and wrote + back carry the internal judged stamp; any re-entry (acompletion + re-entering completion, the async-to-sync /v1/messages dispatch, + interceptor sub-calls reusing the request kwargs) must not re-judge + them, because by then the messages carry litellm's own injected marks + and the judgment would misread those as client breakpoints. + """ + if all(point.get("_litellm_judged") for point in points): + return False + return AnthropicCacheControlHook._request_has_cache_control(messages, system, tools) + @staticmethod def _request_has_cache_control( messages: list[AllMessageValues], @@ -322,7 +353,9 @@ class AnthropicCacheControlHook(CustomPromptManagement): stand down entirely rather than add more, per the auto-caching contract. Tools count: they are a breakpoint the client can mark, they count toward the provider's four-block limit, and caching only the tool definitions is - a common pattern, so injecting alongside them can exceed the cap. + a common pattern, so injecting alongside them can exceed the cap. Tools + carry the mark either at the top level (Anthropic shape) or nested under + ``function`` (OpenAI shape); the Anthropic chat transform accepts both. """ if any(AnthropicCacheControlHook._count_cache_control_blocks(msg) for msg in messages): return True @@ -330,7 +363,14 @@ class AnthropicCacheControlHook(CustomPromptManagement): if any(isinstance(block, dict) and block.get("cache_control") is not None for block in system): return True if tools is not None: - return any(isinstance(tool, dict) and tool.get("cache_control") is not None for tool in tools) + return any( + isinstance(tool, dict) + and ( + tool.get("cache_control") is not None + or (isinstance(tool.get("function"), dict) and tool["function"].get("cache_control") is not None) + ) + for tool in tools + ) return False @staticmethod @@ -392,13 +432,23 @@ class AnthropicCacheControlHook(CustomPromptManagement): custom_llm_provider: str | None, tools: list | None = None, ) -> None: - """For /chat/completions: add default injection points to the request params. + """For /chat/completions: resolve the injection points the request should carry. - No-op when injection points are already configured (explicit config wins). - Seeding the param lets the existing prompt-management gate and the - AnthropicCacheControlHook run unchanged. + Configured injection points win over the automatic defaults, but stand + down entirely when the client already marked its own cache_control + breakpoints (messages or tools): injecting alongside them clashes with + the client's caching strategy and can exceed the provider's four-block + limit. The judgment happens once per request; points a prior pass + wrote back carry the judged stamp and are never re-judged (see + ``_should_stand_down``). Seeding the param lets the existing + prompt-management gate and the AnthropicCacheControlHook run + unchanged. """ if non_default_params.get("cache_control_injection_points"): + if AnthropicCacheControlHook._should_stand_down( + non_default_params["cache_control_injection_points"], messages, None, tools + ): + non_default_params.pop("cache_control_injection_points") return points = AnthropicCacheControlHook.get_default_injection_points( messages=messages, @@ -421,18 +471,26 @@ class AnthropicCacheControlHook(CustomPromptManagement): ) -> Tuple[List[Dict], str | list | None]: """Extract cache_control_injection_points from kwargs and apply if present. - When none are configured but ``litellm.enable_anthropic_prompt_caching`` - is on, synthesize default breakpoints for the native /v1/messages path. - Pops the key from kwargs; if remaining (non-message) points exist they - are written back so downstream transforms can handle them. + Configured points stand down entirely when the client already marked + its own cache_control breakpoints anywhere in the request. The + judgment happens once per request; points a prior pass wrote back + carry the judged stamp and are never re-judged (see + ``_should_stand_down``). When none are configured but + ``litellm.enable_anthropic_prompt_caching`` is on, synthesize default + breakpoints for the native /v1/messages path. Pops the key from kwargs; + if remaining (non-message) points exist they are written back so + downstream transforms can handle them. """ + typed_messages = cast(list[AllMessageValues], messages) # cast-ok: Anthropic-shaped dicts from v1/messages configured = cast( # cast-ok: kwargs is untyped; this key only holds the documented injection-point list list[CacheControlInjectionPoint] | None, kwargs.pop("cache_control_injection_points", None) ) + if configured and AnthropicCacheControlHook._should_stand_down(configured, typed_messages, system, tools): + return messages, system injection_points: list[CacheControlInjectionPoint] = configured or [] if not injection_points and model is not None: injection_points = AnthropicCacheControlHook.get_default_injection_points( - messages=cast(list[AllMessageValues], messages), # cast-ok: Anthropic-shaped dicts from v1/messages + messages=typed_messages, system=system, tools=tools, model=model, @@ -447,7 +505,7 @@ class AnthropicCacheControlHook(CustomPromptManagement): injection_points=injection_points, ) if remaining: - kwargs["cache_control_injection_points"] = remaining + kwargs["cache_control_injection_points"] = AnthropicCacheControlHook._stamped_as_judged(remaining) return messages, system @property diff --git a/litellm/litellm_core_utils/prompt_templates/common_utils.py b/litellm/litellm_core_utils/prompt_templates/common_utils.py index 538d5f650ef..c43089950ee 100644 --- a/litellm/litellm_core_utils/prompt_templates/common_utils.py +++ b/litellm/litellm_core_utils/prompt_templates/common_utils.py @@ -42,6 +42,7 @@ from litellm.types.utils import ( ) if TYPE_CHECKING: # newer pattern to avoid importing pydantic objects on __init__.py + from litellm.types.llms.anthropic import AnthropicInputSchema from litellm.types.llms.openai import ChatCompletionImageObject DEFAULT_USER_CONTINUE_MESSAGE = ChatCompletionUserMessage(content="Please continue.", role="user") @@ -1046,6 +1047,31 @@ def unpack_legacy_defs( return schema +def sanitize_input_schema_for_anthropic(input_schema: dict) -> "AnthropicInputSchema": + """Coerce an arbitrary tool input_schema into the shape Anthropic accepts. + + Anthropic requires ``type == "object"``, only recognises ``$defs`` (legacy + ``definitions`` / OpenAPI ``components.schemas`` refs must be inlined first), + and rejects keys outside ``AnthropicInputSchema``. Both the chat + (``AnthropicConfig._map_tool_helper``) and Anthropic Messages MCP paths run + a schema through here so an external MCP schema cannot succeed on one route + and 400 on the other. + """ + from litellm.types.llms.anthropic import AnthropicInputSchema + + normalized = dict(input_schema) if input_schema else {} + if normalized.get("type") != "object": + normalized["type"] = "object" + if "properties" not in normalized: + normalized["properties"] = {} + + normalized = unpack_legacy_defs(normalized, copy=True) + + allowed_keys = set(AnthropicInputSchema.__annotations__.keys()) + filtered = {key: value for key, value in normalized.items() if key in allowed_keys} + return AnthropicInputSchema(**filtered) + + def _get_image_mime_type_from_url(url: str) -> Optional[str]: """ Get mime type for common image URLs diff --git a/litellm/llms/anthropic/chat/transformation.py b/litellm/llms/anthropic/chat/transformation.py index 0ec1f3eae13..5a0f274e3ca 100644 --- a/litellm/llms/anthropic/chat/transformation.py +++ b/litellm/llms/anthropic/chat/transformation.py @@ -29,7 +29,9 @@ from litellm.constants import ( RESPONSE_FORMAT_TOOL_NAME, ) from litellm.litellm_core_utils.core_helpers import map_finish_reason -from litellm.litellm_core_utils.prompt_templates.common_utils import unpack_legacy_defs +from litellm.litellm_core_utils.prompt_templates.common_utils import ( + sanitize_input_schema_for_anthropic, +) from litellm.llms.base_llm.base_utils import type_to_response_format_param from litellm.llms.base_llm.chat.transformation import BaseConfig, BaseLLMException from litellm.types.llms.anthropic import ( @@ -634,7 +636,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): mcp_server: Optional[AnthropicMcpServerTool] = None if tool["type"] == "function" or tool["type"] == "custom": - _input_schema: dict = tool["function"].get( + _input_schema = tool["function"].get( "parameters", { "type": "object", @@ -642,28 +644,7 @@ class AnthropicConfig(AnthropicModelInfo, BaseConfig): }, ) - # Anthropic requires input_schema.type to be "object". Normalize - # schemas from external sources (MCP servers, OpenAI callers) that - # may omit the type field or use a non-object type. - if _input_schema.get("type") != "object": - litellm.verbose_logger.debug( - "_map_tool_helper: coercing input_schema type from %r to " - "'object' for Anthropic compatibility (tool: %s)", - _input_schema.get("type"), - tool["function"].get("name"), - ) - _input_schema = dict(_input_schema) # avoid mutating caller's dict - _input_schema["type"] = "object" - if "properties" not in _input_schema: - _input_schema["properties"] = {} - - # Inline legacy / OpenAPI $refs before the allow-list filter strips - # their backing def blocks (https://github.com/BerriAI/litellm/issues/26692). - _input_schema = unpack_legacy_defs(_input_schema, copy=True) - - _allowed_properties = set(AnthropicInputSchema.__annotations__.keys()) - input_schema_filtered = {k: v for k, v in _input_schema.items() if k in _allowed_properties} - input_anthropic_schema: AnthropicInputSchema = AnthropicInputSchema(**input_schema_filtered) + input_anthropic_schema = sanitize_input_schema_for_anthropic(_input_schema) _tool = AnthropicMessagesTool( name=tool["function"]["name"], diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py index 703ccf13c27..1a4144de39e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/messages/handler.py +++ b/litellm/llms/anthropic/experimental_pass_through/messages/handler.py @@ -485,6 +485,41 @@ def anthropic_messages_handler( mock_response=litellm_params.mock_response, ) + # Expand litellm_proxy MCP references through the MCP gateway before dispatch, so every + # downstream path (native passthrough and both bridges) gets real tools rather than a + # reference the provider cannot resolve. Popped from kwargs so it never reaches the provider. + skip_mcp_handler = kwargs.pop("_skip_mcp_handler", False) + if not skip_mcp_handler and tools: + from litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler import ( + anthropic_messages_with_mcp, + ) + from litellm.responses.mcp.litellm_proxy_mcp_handler import ( + LiteLLM_Proxy_MCP_Handler, + ) + + if LiteLLM_Proxy_MCP_Handler._should_use_litellm_mcp_gateway(tools=tools): + return anthropic_messages_with_mcp( + max_tokens=max_tokens, + messages=messages, + model=model, + metadata=metadata, + stop_sequences=stop_sequences, + stream=stream, + system=system, + temperature=temperature, + thinking=thinking, + tool_choice=tool_choice, + tools=tools, + top_k=top_k, + top_p=top_p, + container=container, + api_key=api_key, + api_base=api_base, + client=client, + custom_llm_provider=custom_llm_provider, + **kwargs, + ) + anthropic_messages_provider_config: Optional[BaseAnthropicMessagesConfig] = None if custom_llm_provider is not None and custom_llm_provider in [provider.value for provider in LlmProviders]: diff --git a/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py b/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py new file mode 100644 index 00000000000..813d4a62089 --- /dev/null +++ b/litellm/llms/anthropic/experimental_pass_through/messages/mcp_handler.py @@ -0,0 +1,176 @@ +""" +MCP gateway support for the Anthropic `/v1/messages` API. + +Mirrors ``litellm.responses.mcp.chat_completions_handler`` but speaks the +Anthropic Messages shapes: tools carry an ``input_schema``, the model asks for a +tool through a ``tool_use`` content block, and results are fed back as +``tool_result`` blocks in a user message. +""" + +from typing import Any, AsyncIterator, Mapping, Sequence, Union + +from litellm._logging import verbose_logger +from litellm.responses.mcp.request_context import MCPRequestContext +from litellm.types.llms.anthropic import ( + AnthropicMessagesTool, + AnthropicMessagesToolResultParam, + AnthropicMessagesUserMessageParam, +) +from litellm.types.llms.anthropic_messages.anthropic_response import ( + AnthropicMessagesResponse, +) + +MAX_MCP_TOOL_USE_ITERATIONS = 10 + + +def _get_response_content(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]: + content = response.get("content") + if not isinstance(content, list): + return () + return tuple(block for block in content if isinstance(block, dict)) + + +def _extract_tool_use_blocks(response: AnthropicMessagesResponse) -> Sequence[Mapping[str, Any]]: + """Return the ``tool_use`` content blocks the model emitted.""" + return tuple(block for block in _get_response_content(response) if block.get("type") == "tool_use") + + +def _get_stop_reason(response: AnthropicMessagesResponse) -> Union[str, None]: + stop_reason = response.get("stop_reason") + return stop_reason if isinstance(stop_reason, str) else None + + +def _build_tool_result_message(tool_results: Sequence[Mapping[str, Any]]) -> AnthropicMessagesUserMessageParam: + """Turn executed tool results into the user message Anthropic expects.""" + return AnthropicMessagesUserMessageParam( + role="user", + content=tuple( + AnthropicMessagesToolResultParam( + type="tool_result", + tool_use_id=str(result.get("tool_call_id") or ""), + content=str(result.get("result") or ""), + ) + for result in tool_results + ), + ) + + +async def anthropic_messages_with_mcp( + max_tokens: int, + messages: Sequence[Mapping[str, Any]], + model: str, + tools: Union[Sequence[Mapping[str, Any]], None] = None, + **kwargs: Any, # kwargs-ok: forwarded verbatim to litellm.anthropic_messages, which owns the param contract +) -> Union[AnthropicMessagesResponse, AsyncIterator[Any]]: + """ + Expand litellm_proxy MCP references for `/v1/messages` and run the tool loop. + + The MCP gateway owns the expansion so the reference resolves against the + caller's own credentials and access control, rather than being handed to the + upstream provider as a url it cannot reach. + """ + import litellm + from litellm.experimental_mcp_client.tools import ( + transform_mcp_tool_to_anthropic_tool, + ) + from litellm.responses.mcp.litellm_proxy_mcp_handler import ( + LiteLLM_Proxy_MCP_Handler, + ) + + mcp_references, other_tools = LiteLLM_Proxy_MCP_Handler._parse_mcp_tools(tools) + + if not mcp_references: + return await litellm.anthropic_messages( + max_tokens=max_tokens, + messages=list(messages), + model=model, + tools=list(tools) if tools else None, + _skip_mcp_handler=True, + **kwargs, + ) + + context = MCPRequestContext.resolve(kwargs=dict(kwargs), tools=tools) + + ( + deduplicated_mcp_tools, + tool_server_map, + ) = await LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform( + context.user_api_key_auth, + mcp_references, + litellm_trace_id=context.litellm_trace_id, + mcp_auth_header=context.mcp_auth_header, + mcp_server_auth_headers=context.mcp_server_auth_headers, + request_tags=list(context.request_tags) if context.request_tags else None, + ) + + anthropic_tools: Sequence[AnthropicMessagesTool] = tuple( + transform_mcp_tool_to_anthropic_tool(mcp_tool) for mcp_tool in deduplicated_mcp_tools + ) + all_tools = [*anthropic_tools, *(other_tools or ())] + + should_auto_execute = LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools( + mcp_tools_with_litellm_proxy=mcp_references + ) + stream = bool(kwargs.pop("stream", False)) + + base_call_args: Mapping[str, Any] = { + "max_tokens": max_tokens, + "model": model, + "tools": all_tools or None, + "_skip_mcp_handler": True, + **kwargs, + } + + if not should_auto_execute: + return await litellm.anthropic_messages(messages=list(messages), stream=stream, **base_call_args) + + working_messages: Sequence[Mapping[str, Any]] = tuple(messages) + response: AnthropicMessagesResponse = await litellm.anthropic_messages( + messages=list(working_messages), stream=False, **base_call_args + ) + + for _ in range(MAX_MCP_TOOL_USE_ITERATIONS): + if _get_stop_reason(response) != "tool_use": + break + + tool_use_blocks = _extract_tool_use_blocks(response) + if not tool_use_blocks: + break + + tool_results = await LiteLLM_Proxy_MCP_Handler._execute_tool_calls( + tool_server_map=tool_server_map, + tool_calls=list(tool_use_blocks), + user_api_key_auth=context.user_api_key_auth, + mcp_auth_header=context.mcp_auth_header, + mcp_server_auth_headers=context.mcp_server_auth_headers, + oauth2_headers=context.oauth2_headers, + raw_headers=context.raw_headers, + litellm_call_id=context.litellm_call_id, + litellm_trace_id=context.litellm_trace_id, + request_tags=list(context.request_tags) if context.request_tags else None, + ) + + # Every tool call was skipped, so there is nothing to feed back; a + # tool_result message with empty content is rejected by Anthropic. + if not tool_results: + break + + working_messages = ( + *working_messages, + {"role": "assistant", "content": list(_get_response_content(response))}, + _build_tool_result_message(tool_results), + ) + response = await litellm.anthropic_messages(messages=list(working_messages), stream=False, **base_call_args) + else: + verbose_logger.warning( + f"MCP tool loop hit its {MAX_MCP_TOOL_USE_ITERATIONS} iteration cap for model {model}; " + "returning the last response" + ) + + if stream: + from litellm.llms.anthropic.experimental_pass_through.messages.fake_stream_iterator import ( + FakeAnthropicMessagesStreamIterator, + ) + + return FakeAnthropicMessagesStreamIterator(response) + return response diff --git a/litellm/llms/bedrock/audio_transcription/__init__.py b/litellm/llms/bedrock/audio_transcription/__init__.py new file mode 100644 index 00000000000..f2e58df3015 --- /dev/null +++ b/litellm/llms/bedrock/audio_transcription/__init__.py @@ -0,0 +1,84 @@ +import base64 +from typing import Union + +import httpx + +from litellm.litellm_core_utils.audio_utils.utils import process_audio_file +from litellm.rust_bridge import transcription as rust_transcription_bridge +from litellm.types.utils import FileTypes, TranscriptionResponse + + +class BedrockAudioTranscriptionRustDispatch: + @staticmethod + def _audio_payload(audio_file: FileTypes) -> dict[str, object]: + processed_audio = process_audio_file(audio_file) + formats = { + "audio/flac": "flac", + "audio/mpeg": "mp3", + "audio/mp3": "mp3", + "audio/ogg": "ogg", + "audio/wav": "wav", + "audio/x-wav": "wav", + } + audio_format = formats.get(processed_audio.content_type) or ( + processed_audio.filename.rsplit(".", 1)[-1].lower() if "." in processed_audio.filename else "" + ) + if audio_format not in {"wav", "mp3", "flac", "ogg"}: + raise ValueError(f"Unsupported Bedrock audio format for file {processed_audio.filename!r}") + return { + "data": base64.b64encode(processed_audio.file_content).decode("ascii"), + "format": audio_format, + "filename": processed_audio.filename, + } + + def audio_transcriptions( + self, + *, + model: str, + audio_file: FileTypes, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout: Union[float, httpx.Timeout] | None, + ) -> TranscriptionResponse: + rust_response = rust_transcription_bridge.transcription( + model=model, + audio=self._audio_payload(audio_file), + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ) + if rust_response is None: + raise RuntimeError("Rust audio transcription bridge is unavailable") + return TranscriptionResponse(**rust_response) + + async def async_audio_transcriptions( + self, + *, + model: str, + audio_file: FileTypes, + api_key: str | None, + api_base: str | None, + custom_llm_provider: str, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout: Union[float, httpx.Timeout] | None, + ) -> TranscriptionResponse: + rust_response = await rust_transcription_bridge.atranscription( + model=model, + audio=self._audio_payload(audio_file), + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ) + if rust_response is None: + raise RuntimeError("Rust audio transcription bridge is unavailable") + return TranscriptionResponse(**rust_response) diff --git a/litellm/llms/bedrock/batches/transformation.py b/litellm/llms/bedrock/batches/transformation.py index 4fcf7cf91cb..a4ff1c78467 100644 --- a/litellm/llms/bedrock/batches/transformation.py +++ b/litellm/llms/bedrock/batches/transformation.py @@ -4,6 +4,7 @@ import time from typing import Any, Dict, List, Literal, Optional, Union, cast from httpx import Headers, Response +from pydantic import TypeAdapter, ValidationError from litellm.litellm_core_utils.cloud_storage_security import ( BEDROCK_MANAGED_S3_BATCH_PREFIX, @@ -19,6 +20,7 @@ from litellm.types.llms.bedrock import ( BedrockOutputDataConfig, BedrockS3InputDataConfig, BedrockS3OutputDataConfig, + BedrockTag, ) from litellm.types.llms.openai import ( AllMessageValues, @@ -38,6 +40,18 @@ _S3_BATCH_FILE_UUID_SUFFIX_PATTERN = re.compile( r"-[0-9a-fA-F]{8}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{4}-[0-9a-fA-F]{12}\.jsonl$" ) +_BEDROCK_TAGS_ADAPTER: TypeAdapter[list[BedrockTag]] = TypeAdapter(list[BedrockTag]) + + +def _validate_bedrock_tags(raw_tags: object) -> list[BedrockTag]: + try: + return _BEDROCK_TAGS_ADAPTER.validate_python(raw_tags, strict=True) + except ValidationError as e: + raise ValueError( + "Invalid 'bedrock_tags' value. Expected a list of {'key': , 'value': } dicts, " + f"e.g. [{{'key': 'team', 'value': 'genai'}}]. Got: {raw_tags!r}" + ) from e + class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): """ @@ -201,6 +215,11 @@ class BedrockBatchesConfig(BaseAWSLLM, BaseBatchesConfig): "roleArn": role_arn, } + config_bedrock_tags = litellm_params.get("bedrock_tags") + bedrock_tags = config_bedrock_tags if config_bedrock_tags is not None else optional_params.get("bedrock_tags") + if bedrock_tags is not None: + bedrock_request["tags"] = _validate_bedrock_tags(bedrock_tags) + # Add optional parameters if provided completion_window = create_batch_data.get("completion_window") if completion_window: diff --git a/litellm/llms/fireworks_ai/chat/transformation.py b/litellm/llms/fireworks_ai/chat/transformation.py index 319f03fea89..eeae8c76888 100644 --- a/litellm/llms/fireworks_ai/chat/transformation.py +++ b/litellm/llms/fireworks_ai/chat/transformation.py @@ -133,6 +133,32 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig): def get_config(cls): return super().get_config() + def validate_environment( + self, + headers: dict, + model: str, + messages: List[AllMessageValues], + optional_params: dict, + litellm_params: dict, + api_key: str | None = None, + api_base: str | None = None, + ) -> dict: + api_key = self._get_api_key(api_key) + if api_key is None: + raise ValueError("FIREWORKS_API_KEY is not set") + + validated_headers = OpenAIGPTConfig.validate_environment( + self, + headers=headers, + model=model, + messages=messages, + optional_params=optional_params, + litellm_params=litellm_params, + api_key=api_key, + api_base=api_base, + ) + return self._add_session_affinity_header(validated_headers, litellm_params) + def get_supported_openai_params(self, model: str): # Base parameters supported by all models supported_params = [ diff --git a/litellm/llms/fireworks_ai/common_utils.py b/litellm/llms/fireworks_ai/common_utils.py index 4e22445bcc0..51ed8afbbd2 100644 --- a/litellm/llms/fireworks_ai/common_utils.py +++ b/litellm/llms/fireworks_ai/common_utils.py @@ -64,9 +64,16 @@ class FireworksAIMixin: if api_key is None: raise ValueError("FIREWORKS_API_KEY is not set") - validated_headers = {"Authorization": "Bearer {}".format(api_key), **headers} - if not any(key.lower() == "x-session-affinity" for key in validated_headers): - session_id = get_fireworks_session_id(litellm_params) - if session_id: - validated_headers["x-session-affinity"] = session_id - return validated_headers + auth_headers = {"Authorization": "Bearer {}".format(api_key), **headers} + content_type_header = ( + {} if any(key.lower() == "content-type" for key in auth_headers) else {"Content-Type": "application/json"} + ) + return self._add_session_affinity_header({**auth_headers, **content_type_header}, litellm_params) + + def _add_session_affinity_header(self, headers: dict, litellm_params: dict) -> dict: + if any(key.lower() == "x-session-affinity" for key in headers): + return headers + session_id = get_fireworks_session_id(litellm_params) + if not session_id: + return headers + return {**headers, "x-session-affinity": session_id} diff --git a/litellm/main.py b/litellm/main.py index 3584297b35f..fb05a375111 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -27,7 +27,6 @@ from typing import ( TYPE_CHECKING, Any, AsyncIterator, - Callable, Coroutine, Dict, Iterable, @@ -81,22 +80,19 @@ from litellm.constants import ( from litellm.exceptions import LiteLLMUnknownProvider from litellm.integrations.custom_logger import CustomLogger from litellm.litellm_core_utils.asyncify import run_async_function -from litellm.litellm_core_utils.chat_completion_agentic_loop import ( - maybe_run_chat_completion_agentic_loop, -) from litellm.litellm_core_utils.audio_utils.utils import ( calculate_request_duration, get_audio_file_for_health_check, ) -from litellm.litellm_core_utils.completion_timeout import CompletionTimeout -from litellm.litellm_core_utils.request_timeout_resolver import ( - get_configured_request_timeout, +from litellm.litellm_core_utils.chat_completion_agentic_loop import ( + maybe_run_chat_completion_agentic_loop, ) +from litellm.litellm_core_utils.completion_timeout import CompletionTimeout +from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_litellm_params import ( AWS_CREDENTIAL_KWARGS_KEYS, OPTIONAL_KWARGS_KEYS, ) -from litellm.litellm_core_utils.dd_tracing import tracer from litellm.litellm_core_utils.get_provider_specific_headers import ( ProviderSpecificHeaderUtils, ) @@ -112,6 +108,9 @@ from litellm.litellm_core_utils.mock_functions import ( from litellm.litellm_core_utils.prompt_templates.common_utils import ( get_content_from_model_response, ) +from litellm.litellm_core_utils.request_timeout_resolver import ( + get_configured_request_timeout, +) from litellm.llms.base_llm import BaseConfig, BaseImageGenerationConfig from litellm.llms.base_llm.base_model_iterator import ( convert_model_response_to_streaming, @@ -213,7 +212,6 @@ from .llms.bedrock.embed.embedding import BedrockEmbedding from .llms.bedrock.image_edit.handler import BedrockImageEdit from .llms.bedrock.image_generation.image_handler import BedrockImageGeneration from .llms.bytez.chat.transformation import BytezChatConfig -from .llms.gdc.chat.transformation import GDCGeminiConfig from .llms.clarifai.chat.transformation import ClarifaiConfig from .llms.codestral.completion.handler import CodestralTextCompletion from .llms.cohere.embed import handler as cohere_embed @@ -222,24 +220,25 @@ from .llms.custom_httpx.llm_http_handler import BaseLLMHTTPHandler from .llms.custom_llm import CustomLLM, custom_chat_llm_router from .llms.databricks.embed.handler import DatabricksEmbeddingHandler from .llms.deprecated_providers import aleph_alpha, palm +from .llms.gdc.chat.transformation import GDCGeminiConfig from .llms.gemini.common_utils import get_api_key_from_env from .llms.groq.chat.handler import GroqChatCompletion from .llms.heroku.chat.transformation import HerokuChatConfig from .llms.huggingface.embedding.handler import HuggingFaceEmbedding from .llms.lemonade.chat.transformation import LemonadeChatConfig from .llms.nlp_cloud.chat.handler import completion as nlp_cloud_chat_completion -from .llms.oci.chat.transformation import OCIChatConfig -from .llms.ollama.completion import handler as ollama -from .llms.oobabooga.chat import oobabooga -from .llms.openai.completion.handler import OpenAITextCompletion -from .llms.openai.image_variations.handler import OpenAIImageVariationsHandler -from .llms.openai.openai import OpenAIChatCompletion from .llms.nvidia_riva.audio_transcription.handler import ( NvidiaRivaAudioTranscription, ) from .llms.nvidia_riva.audio_transcription.transformation import ( NvidiaRivaAudioTranscriptionConfig, ) +from .llms.oci.chat.transformation import OCIChatConfig +from .llms.ollama.completion import handler as ollama +from .llms.oobabooga.chat import oobabooga +from .llms.openai.completion.handler import OpenAITextCompletion +from .llms.openai.image_variations.handler import OpenAIImageVariationsHandler +from .llms.openai.openai import OpenAIChatCompletion from .llms.openai.transcriptions.handler import OpenAIAudioTranscription from .llms.openai_like.chat.handler import OpenAILikeChatHandler from .llms.openai_like.embedding.handler import OpenAILikeEmbeddingHandler @@ -7722,6 +7721,32 @@ def transcription( headers=extra_headers, provider_config=provider_config, # type: ignore[arg-type] ) + elif custom_llm_provider == "bedrock": + from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch + + dispatch = BedrockAudioTranscriptionRustDispatch() + if atranscription: + response = dispatch.async_audio_transcriptions( + model=model, + audio_file=file, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ) + else: + response = dispatch.audio_transcriptions( + model=model, + audio_file=file, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout=timeout, + ) elif provider_config is not None: response = base_llm_http_handler.audio_transcriptions( model=model, diff --git a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py index 1ba608b9510..d06dad34d8c 100644 --- a/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py +++ b/litellm/proxy/_experimental/mcp_server/mcp_server_manager.py @@ -93,6 +93,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.token_exchange_ ) from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( AuthorizationCodeConfig, + ClientCredentialsConfig, CredError, IdJagConfig, PassthroughConfig, @@ -2739,14 +2740,18 @@ class MCPServerManager: ) if not conflicts: return auth, extra_headers - if isinstance(spec.config, (TokenExchangeConfig, AuthorizationCodeConfig, IdJagConfig)): - # The resolver owns the per-user credential here (token_exchange's exchanged - # token, authorization_code's stored token, id_jag's minted assertion). It is - # authoritative: a guardrail such - # as MCPJWTSigner, static_headers, or any other injected Authorization must NOT - # shadow it (otherwise the upstream gets e.g. the signer's JWT instead of the - # exchanged token and rejects it). Drop the conflicting header so the resolved - # token reaches upstream. + if isinstance( + spec.config, + (TokenExchangeConfig, AuthorizationCodeConfig, IdJagConfig, ClientCredentialsConfig), + ): + # The resolver owns the credential here (token_exchange's exchanged token, + # authorization_code's stored token, id_jag's minted assertion, + # client_credentials' gateway-minted M2M token). It is authoritative: a + # guardrail such as MCPJWTSigner, static_headers, or any other injected + # Authorization must NOT shadow it (otherwise the upstream gets e.g. the + # signer's JWT instead of the minted token and rejects it, and for M2M the + # one-shot 401 refetch is lost with it). Drop the conflicting header so the + # resolved token reaches upstream. return auth, _without_authorization(extra_headers) # Other modes: an Authorization already supplied via extra_headers (a forwarded caller # header or static_headers) is intentional and wins; v1 applies those last. diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py index 6631e38f524..565c489e77c 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/adapter.py @@ -22,6 +22,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( ApiKeyConfig, AuthorizationCodeConfig, ClientAuth, + ClientCredentialsConfig, ClientSecretAuth, CredError, IdJagConfig, @@ -70,10 +71,10 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]: an ``assert_never`` tail, so a newly added auth mode fails the type gate here until it is explicitly mapped or explicitly deferred, rather than silently falling through to v1. Live modes: ``none``, the static-header family (``api_key`` plus the Authorization schemes, - all shared-key), ``oauth2`` per-user tokens (``authorization_code``), ``oauth2_token_exchange`` - (OBO), and the client-forwarded token modes ``true_passthrough`` / ``oauth_delegate`` - (``PassthroughConfig``); client_credentials (M2M), delegated/passthrough oauth2, and SigV4 - return None and stay on v1. + all shared-key), ``oauth2`` per-user tokens (``authorization_code``), ``oauth2`` M2M + (``client_credentials``), ``oauth2_token_exchange`` (OBO), and the client-forwarded token + modes ``true_passthrough`` / ``oauth_delegate`` (``PassthroughConfig``); delegated/passthrough + oauth2 and SigV4 return None and stay on v1. """ if server.is_byok: return None # per-user BYOK source not migrated yet -> defer to v1 (any auth_type) @@ -95,14 +96,7 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]: case MCPAuth.basic: return _shared_key_spec(server, resource, "Authorization", "Basic", encode=True) case MCPAuth.oauth2: - if server.needs_user_oauth_token and not server.delegate_auth_to_upstream: - return ServerSpec( - server_id=server.server_id, - resource=resource, - config=AuthorizationCodeConfig(), - ) - # client_credentials (M2M) and delegate/passthrough oauth2 stay on v1 - return None + return _oauth2_spec(server, resource) case MCPAuth.oauth2_id_jag: return _id_jag_spec(server, resource) case MCPAuth.true_passthrough | MCPAuth.oauth_delegate: @@ -114,6 +108,47 @@ def to_server_spec(server: MCPServer) -> Optional[ServerSpec]: assert_never(auth_type) +def _oauth2_spec(server: MCPServer, resource: str) -> ServerSpec | None: + """Dispatch the oauth2 auth_type across its sub-modes: M2M, gateway-managed interactive, or v1. + + ``client_credentials`` (the explicit ``oauth2_flow`` opt-in) builds the M2M spec, per-user + ``authorization_code`` without upstream delegation builds the interactive spec, and the + delegate/passthrough shapes defer to v1 (None). + """ + if server.has_client_credentials: + return _client_credentials_spec(server, resource) + if server.needs_user_oauth_token and not server.delegate_auth_to_upstream: + return ServerSpec( + server_id=server.server_id, + resource=resource, + config=AuthorizationCodeConfig(), + ) + return None + + +def _client_credentials_spec(server: MCPServer, resource: str) -> ServerSpec: + """Build a client_credentials (M2M) spec; the explicit ``oauth2_flow`` opt-in owns the server. + + Missing grant fields (``client_id``/``client_secret``/``token_url``) are NOT a reason to defer: + v1 would connect unauthenticated and the upstream's 401 gets absorbed into an empty tool list, + so the arm fails closed with ``misconfigured`` instead, naming the missing fields (mirrors the + OBO ownership rule). ``audience`` is forwarded only when the operator set it; a missing one is + omitted, not derived, since a fabricated value risks the IdP rejecting the grant. + """ + return ServerSpec( + server_id=server.server_id, + resource=resource, + config=ClientCredentialsConfig( + client_id=server.client_id, + client_secret=SecretStr(server.client_secret) if server.client_secret else None, + token_url=server.token_url, + scopes=tuple(server.scopes or ()), + audience=server.audience, + token_endpoint_auth_method=server.token_endpoint_auth_method, + ), + ) + + def _token_exchange_spec(server: MCPServer, resource: str) -> Optional[ServerSpec]: """Build a token_exchange (OBO) spec, or defer (None) when it is not OBO-configured. diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py new file mode 100644 index 00000000000..9be1121126a --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/client_credentials.py @@ -0,0 +1,348 @@ +"""The ``client_credentials`` (M2M) arm's token source and retrying bearer auth. + +Implements the client-credentials behavior contract for the v2 resolver: + +- **Acquisition**: POST ``grant_type=client_credentials`` to the configured token endpoint with + the configured scopes and (when set) the IdP's ``audience`` parameter, authenticating the + client per ``token_endpoint_auth_method`` (RFC 6749 section 2.3.1, shared helper). +- **Caching**: tokens are cached per ``(client identity, server)`` where the identity key hashes + ``token_url`` / ``client_id`` / ``client_secret`` / auth method / scopes / audience — rotating + or re-scoping the credentials changes the key, so a stale token can never be served for the + new identity (the contract's rotation-invalidation clause). +- **Expiry**: the cache TTL respects ``expires_in`` minus a skew so an entry lapses before the + real token does; a response with no ``expires_in`` is cached briefly + (``default_ttl_seconds``), not assumed long-lived. No refresh_token is ever expected. +- **401 recovery**: ``ClientCredentialsBearerAuth`` retries an upstream request exactly once + after a 401 — discard the cached token, mint a fresh one, resend; a second failure surfaces + the upstream's own auth error unchanged. +- **No user context**: nothing here reads a ``Subject``; every caller shares the one client + identity. + +The token-endpoint POST is injected (``M2MTokenEndpointPost``) so the grant orchestration is +testable without a live IdP; ``post_client_credentials_grant`` is the httpx edge and the one +place the untyped response boundary is contained. Failures are values: the source returns +``Result[OAuthToken, CredError]``; only the httpx edge touches exceptions. +""" + +from __future__ import annotations + +import asyncio +import hashlib +import time +from collections.abc import AsyncGenerator, Awaitable, Callable, Generator +from dataclasses import dataclass +from typing import Annotated, Literal + +import httpx +from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError +from typing_extensions import assert_never + +from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + InMemoryTokenCacheBackend, + OAuthToken, + TokenCacheBackend, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.result import ( + Error, + Ok, + Result, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + ClientCredentialsConfig, + CredError, +) + + +class TokenEndpointSuccess(BaseModel): + """The endpoint returned a JSON object; field validation is the caller's job.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["success"] = "success" + body: dict[str, object] + + +class TokenEndpointDenied(BaseModel): + """The endpoint answered but did not grant a token (an HTTP error or a non-JSON body).""" + + model_config = ConfigDict(frozen=True) + tag: Literal["denied"] = "denied" + status_code: int + detail: str + + +class TokenEndpointUnreachable(BaseModel): + """The endpoint could not be reached (DNS, TLS, connect/read failure).""" + + model_config = ConfigDict(frozen=True) + tag: Literal["unreachable"] = "unreachable" + detail: str + + +TokenEndpointOutcome = Annotated[ + TokenEndpointSuccess | TokenEndpointDenied | TokenEndpointUnreachable, + Field(discriminator="tag"), +] + +M2MTokenEndpointPost = Callable[[str, "dict[str, str]", "dict[str, str]"], Awaitable[TokenEndpointOutcome]] + + +_TOKEN_BODY_ADAPTER: TypeAdapter[dict[str, object]] = TypeAdapter(dict[str, object]) + + +async def post_client_credentials_grant( + url: str, form: dict[str, str], headers: dict[str, str] +) -> TokenEndpointOutcome: + """POST the grant to the token endpoint and classify the transport outcome. + + The httpx edge: litellm's handler is partially typed (and raises ``HTTPStatusError`` itself on + a 4xx/5xx), so the untyped boundary is contained here and every field the caller reads comes + out of a validated ``TokenEndpointOutcome``. + """ + 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 is partially typed + ) + from litellm.types.llms.custom_http import httpxSpecialProvider # noqa: PLC0415 # deferred with the handler import + + try: + client = get_async_httpx_client(llm_provider=httpxSpecialProvider.Oauth2Check) + response = await client.post( # pyright: ignore[reportUnknownMemberType, reportUnknownVariableType] # handler is partially typed + url, headers={"Accept": "application/json", **headers}, data=form + ) + except httpx.HTTPStatusError as status_err: + status_code = status_err.response.status_code + 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)) + if not isinstance(response, httpx.Response): + return TokenEndpointUnreachable(detail="token endpoint returned no response") + try: + body = _TOKEN_BODY_ADAPTER.validate_json(response.content) + except ValidationError: + return TokenEndpointDenied( + status_code=response.status_code, detail="token endpoint returned a non-JSON-object body" + ) + return TokenEndpointSuccess(body=body) + + +def _parse_expires_in(raw: object) -> int | None: + if isinstance(raw, bool): + return None + if isinstance(raw, int): + return raw + if isinstance(raw, str): + try: + return int(raw) + except ValueError: + return None + return None + + +def _parse_granted_scopes(raw: object) -> tuple[str, ...] | None: + return tuple(raw.split()) if isinstance(raw, str) and raw else None + + +@dataclass(frozen=True, slots=True) +class _PreparedGrant: + """A validated, ready-to-POST grant plus the identity key its token caches under.""" + + token_url: str + form: dict[str, str] + headers: dict[str, str] + identity_key: str + + +class ClientCredentialsTokenSource: + """Cached M2M access tokens, one per ``(client identity, server)``. + + ``get`` serves from the cache while the entry's TTL (derived from ``expires_in`` minus + ``expiry_skew_seconds``) holds, fetching under a per-server lock so concurrent misses + produce one grant. ``refetch`` is the 401-recovery path: it drops the failed token and + mints a fresh one, unless a concurrent caller already replaced it. + """ + + def __init__( + self, + post: M2MTokenEndpointPost = post_client_credentials_grant, + *, + backend: TokenCacheBackend | None = None, + default_ttl_seconds: float = 300.0, + expiry_skew_seconds: float = 60.0, + min_cache_seconds: float = 10.0, + max_locks: int = 1024, + clock: Callable[[], float] = time.time, + ) -> None: + self._post = post + self._backend: TokenCacheBackend = backend or InMemoryTokenCacheBackend(clock=clock) + self._default_ttl_seconds = default_ttl_seconds + self._expiry_skew_seconds = expiry_skew_seconds + self._min_cache_seconds = min_cache_seconds + self._max_locks = max_locks + self._clock = clock + self._locks: dict[str, asyncio.Lock] = {} + + def _lock(self, server_id: str) -> asyncio.Lock: + """Per-server single-flight lock, bounded so ephemeral server ids (e.g. the REST tools + preview mints a fresh id per call) cannot grow the dict for the life of the process. + Evicting the oldest entry while a task still holds it only means a concurrent caller for + that server may run its own grant — single-flight is an optimization, not correctness. + """ + if server_id not in self._locks and len(self._locks) >= self._max_locks: + self._locks.pop(next(iter(self._locks)), None) + return self._locks.setdefault(server_id, asyncio.Lock()) + + async def get(self, server_id: str, config: ClientCredentialsConfig) -> Result[OAuthToken, CredError]: + match _prepare_grant(config): + case Error(err): + return Error(err) + case Ok(grant): + cached = await self._backend.get(grant.identity_key, server_id) + if cached is not None: + return Ok(cached) + async with self._lock(server_id): + cached = await self._backend.get(grant.identity_key, server_id) + if cached is not None: + return Ok(cached) + return await self._fetch_and_cache(server_id, grant) + + async def refetch(self, server_id: str, config: ClientCredentialsConfig, failed_access_token: str) -> str | None: + """Replace a token the upstream just 401'd; returns the fresh bearer value or ``None``. + + Runs under the same per-server lock as ``get``: if a concurrent caller already replaced + the failed token, that replacement is returned without another grant, so a burst of 401s + yields one fetch. A failed refetch returns ``None`` and the caller surfaces the + upstream's original auth error (the contract's retry-once-then-give-up clause). + """ + match _prepare_grant(config): + case Error(_): + return None + case Ok(grant): + async with self._lock(server_id): + cached = await self._backend.get(grant.identity_key, server_id) + if cached is not None and cached.access_token != failed_access_token: + return cached.access_token + await self._backend.delete(grant.identity_key, server_id) + match await self._fetch_and_cache(server_id, grant): + case Ok(token): + return token.access_token + case Error(_): + return None + + async def _fetch_and_cache(self, server_id: str, grant: _PreparedGrant) -> Result[OAuthToken, CredError]: + outcome = await self._post(grant.token_url, grant.form, grant.headers) + match outcome: + case TokenEndpointUnreachable(): + return Error(CredError.of_upstream_unavailable(f"OAuth2 token endpoint unreachable: {outcome.detail}")) + case TokenEndpointDenied(): + if outcome.status_code >= 500: + return Error(CredError.of_upstream_unavailable(f"OAuth2 token endpoint failed: {outcome.detail}")) + return Error(CredError.of_misconfigured(f"OAuth2 client_credentials grant rejected: {outcome.detail}")) + case TokenEndpointSuccess(): + return await self._cache_token(server_id, grant, outcome.body) + assert_never(outcome) + + async def _cache_token( + self, server_id: str, grant: _PreparedGrant, body: dict[str, object] + ) -> Result[OAuthToken, CredError]: + access_token = body.get("access_token") + if not isinstance(access_token, str) or not access_token: + return Error(CredError.of_misconfigured("OAuth2 token response is missing 'access_token'")) + expires_in = _parse_expires_in(body.get("expires_in")) + token = OAuthToken( + access_token=access_token, + expires_at=self._clock() + expires_in if expires_in is not None else None, + scopes=_parse_granted_scopes(body.get("scope")) or (), + ) + # The min-cache floor is itself capped at the token's real lifetime, so a token whose + # expires_in is below the skew is never served past its actual expiry; a non-positive + # expires_in caches nothing (every request re-fetches, serialized by the per-server lock). + ttl = ( + max(expires_in - self._expiry_skew_seconds, min(float(expires_in), self._min_cache_seconds), 0.0) + if expires_in is not None + else self._default_ttl_seconds + ) + if ttl > 0: + await self._backend.set(grant.identity_key, server_id, token, ttl) + return Ok(token) + + +def _prepare_grant(config: ClientCredentialsConfig) -> Result[_PreparedGrant, CredError]: + if not config.client_id or not config.client_secret or not config.token_url: + missing = ", ".join( + name + for name, present in ( + ("client_id", bool(config.client_id)), + ("client_secret", bool(config.client_secret)), + ("token_url", bool(config.token_url)), + ) + if not present + ) + return Error(CredError.of_misconfigured(f"client_credentials config is missing: {missing}")) + + from litellm.proxy._experimental.mcp_server.auth.token_endpoint_auth import ( # noqa: PLC0415 # keep package v1-free at import time + build_token_endpoint_client_auth, + ) + + client_auth = build_token_endpoint_client_auth( + auth_method=config.token_endpoint_auth_method, + client_id=config.client_id, + client_secret=config.client_secret.get_secret_value(), + ) + form = { + "grant_type": "client_credentials", + **client_auth.body, + **({"scope": " ".join(config.scopes)} if config.scopes else {}), + **({"audience": config.audience} if config.audience else {}), + } + return Ok( + _PreparedGrant( + token_url=config.token_url, + form=form, + headers=client_auth.headers, + identity_key=_identity_key(config), + ) + ) + + +def _identity_key(config: ClientCredentialsConfig) -> str: + """Hash of everything that names the client identity; any rotation yields a new key.""" + material = "\n".join( + ( + config.token_url or "", + config.client_id or "", + config.client_secret.get_secret_value() if config.client_secret else "", + config.token_endpoint_auth_method or "", + " ".join(config.scopes), + config.audience or "", + ) + ) + return hashlib.sha256(material.encode("utf-8")).hexdigest() + + +class ClientCredentialsBearerAuth(httpx.Auth): + """Bearer auth that retries an upstream 401 exactly once with a freshly minted token. + + The initial token was already resolved (so config/IdP failures surfaced as typed errors + before any upstream request); ``refetch`` is the source's 401-recovery callback. If the + refetch fails, or the retried request 401s again, the upstream's response stands. + """ + + def __init__(self, access_token: str, refetch: Callable[[str], Awaitable[str | None]]) -> None: + self.header_name = "Authorization" + self._access_token = SecretStr(access_token) + self._refetch = refetch + + async def async_auth_flow(self, request: httpx.Request) -> AsyncGenerator[httpx.Request, httpx.Response]: + token = self._access_token.get_secret_value() + request.headers[self.header_name] = f"Bearer {token}" + response = yield request + if response.status_code != 401: + return + fresh = await self._refetch(token) + if fresh is None: + return + self._access_token = SecretStr(fresh) + request.headers[self.header_name] = f"Bearer {fresh}" + yield request + + def sync_auth_flow(self, request: httpx.Request) -> Generator[httpx.Request, httpx.Response, None]: + raise RuntimeError("ClientCredentialsBearerAuth only supports async httpx clients") diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py index 7e5c073870a..69984a56311 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/resolver.py @@ -9,18 +9,24 @@ at runtime instead of returning `None`. `none`, `api_key` (shared-key source), and `passthrough` (forwards the caller's own inbound token) are live, as is `authorization_code`, which reads the user's token from the injected -`OAuthTokenStore`, and `token_exchange`, which swaps the caller's inbound token through the -injected `TokenExchanger`. The remaining arms are `not_implemented` stubs that each land in a -follow-up PR with their seam. Pure v2: no imports from v1. +`OAuthTokenStore`, `token_exchange`, which swaps the caller's inbound token through the injected +`TokenExchanger`, and `client_credentials`, which mints and caches the gateway's M2M token through +the injected `ClientCredentialsTokenSource`. The remaining arms are `not_implemented` stubs that +each land in a follow-up PR with their seam. Pure v2: no imports from v1. """ from __future__ import annotations import hashlib +from functools import partial import httpx from typing_extensions import assert_never +from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import ( + ClientCredentialsBearerAuth, + ClientCredentialsTokenSource, +) from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( NoOpAuth, StaticHeaderAuth, @@ -104,11 +110,13 @@ class UpstreamCredentialProvider: token_exchanger: TokenExchanger | None = None, token_endpoint: TokenEndpointClient | None = None, exchanged_tokens: ExchangedTokenCache | None = None, + client_credentials_source: ClientCredentialsTokenSource | None = None, ) -> None: self._oauth_token_store: OAuthTokenStore = oauth_token_store or _NullOAuthTokenStore() self._token_exchanger: TokenExchanger = token_exchanger or _NullTokenExchanger() self._token_endpoint: TokenEndpointClient = token_endpoint or TokenEndpointClient() self._exchanged_tokens: ExchangedTokenCache = exchanged_tokens or ExchangedTokenCache() + self._client_credentials_source = client_credentials_source or ClientCredentialsTokenSource() async def resolve_credentials(self, subject: Subject, server: ServerSpec) -> Result[httpx.Auth, CredError]: match server.config: @@ -118,8 +126,8 @@ class UpstreamCredentialProvider: return self._api_key(config) case PassthroughConfig(): return self._passthrough(subject) - case ClientCredentialsConfig(): - return _not_implemented(AuthSpecKind.client_credentials) + case ClientCredentialsConfig() as config: + return await self._client_credentials(server.server_id, config) case TokenExchangeConfig() as config: return await self._token_exchange(subject, server, config) case IdJagConfig() as config: @@ -215,6 +223,23 @@ class UpstreamCredentialProvider: return Error(CredError.of_unauthorized("Authorization required: complete the OAuth flow for this server.")) return Ok(StaticHeaderAuth(f"Bearer {token.access_token}", header_name="Authorization")) + async def _client_credentials( + self, server_id: str, config: ClientCredentialsConfig + ) -> Result[httpx.Auth, CredError]: + """The M2M arm: resolve a cached (or freshly minted) gateway token; no user context. + + The token is resolved here, before any upstream request, so a misconfigured grant or an + unreachable IdP surfaces as a typed ``CredError``. The returned auth carries the source's + ``refetch``, so an upstream 401 is retried exactly once with a freshly minted token (the + contract's invalid-token recovery); a second 401 surfaces the upstream's own error. + """ + match await self._client_credentials_source.get(server_id, config): + case Ok(token): + refetch = partial(self._client_credentials_source.refetch, server_id, config) + return Ok(ClientCredentialsBearerAuth(token.access_token, refetch)) + case Error(err): + return Error(err) + async def _token_exchange( self, subject: Subject, server: ServerSpec, config: TokenExchangeConfig ) -> Result[StaticHeaderAuth, CredError]: @@ -245,7 +270,9 @@ class UpstreamCredentialProvider: Used after an upstream rejects the injected credential, so the next resolve re-mints rather than serving the same rejected token until TTL. `token_exchange` and `id_jag` hold a - re-mintable cached credential here; other modes are a no-op. + re-mintable cached credential here; `client_credentials` recovers inside its own auth flow + (`ClientCredentialsBearerAuth` retries the 401'd request once with a fresh token), and + other modes are a no-op. """ if subject.inbound_token is None: return diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py new file mode 100644 index 00000000000..08d5cc8b1f1 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_credentials.py @@ -0,0 +1,190 @@ +"""Producer and consumer helpers for the gateway-level DCR session token. + +The aggregate ``/mcp`` front door (``mcp_gateway_dcr``) issues the identity-only session +tokens defined in :mod:`.session_token`. The gateway token endpoint mints them (producer) +after SSO sign-in, and at the MCP admission edge the gateway derives the session signing +key from the proxy ``master_key``, opens the bearer, and admits the request under the +recovered litellm user (consumer), reloading the live user record and policy before +anything runs. This module is the pure surface for both sides; the token-endpoint and +admission wiring live in their respective call sites. + +The signing key is derived with the same memory-hard scrypt construction as +:func:`~.bridge_credentials.envelope_keys_from_master_key` but under a distinct domain +label, so session tokens and bridge envelopes never share key material: a token of one +family is unverifiable in the other by key separation, on top of the distinct issuers, +prefixes, and claim shapes. +""" + +import hashlib +from datetime import datetime +from functools import lru_cache +from typing import Literal, TypeAlias + +from pydantic import BaseModel, ConfigDict, SecretStr + +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + OpenedSessionToken, + SessionExpired, + SessionKeys, + SessionPrincipal, + is_session_refresh_token, + is_session_token, + open_session_refresh_token, + open_session_token, +) + +_SESSION_SIGNING_KEY_DOMAIN = b"litellm-mcp-gateway:session-signing:" + +# scrypt work factors (RFC 7914), identical to the envelope KDF: memory-hard so a captured +# session token is not a cheap offline oracle for the master key. +_SCRYPT_N = 2**15 +_SCRYPT_R = 8 +_SCRYPT_P = 1 +_SCRYPT_MAXMEM = 128 * _SCRYPT_N * _SCRYPT_R * _SCRYPT_P * 2 +_DERIVED_KEY_BYTES = 32 + + +@lru_cache(maxsize=8) +def session_keys_from_master_key(master_key: str) -> SessionKeys: + """Derive the session signing key from the proxy master key. + + A memory-hard scrypt KDF (RFC 7914) over a session-specific domain-label salt yields a + 256-bit subkey from the one secret, so the producer (mint) and consumer (open) agree on + the key without persisting any. The domain label differs from both envelope labels in + :mod:`.bridge_credentials`, so compromise or misuse of one token family never crosses + into the other. The result is cached (the master key is fixed for a process); rotating + ``master_key`` invalidates every outstanding session, which is the intended behavior + for a signing-key change. + """ + signing = hashlib.scrypt( + master_key.encode(), + salt=_SESSION_SIGNING_KEY_DOMAIN, + n=_SCRYPT_N, + r=_SCRYPT_R, + p=_SCRYPT_P, + maxmem=_SCRYPT_MAXMEM, + dklen=_DERIVED_KEY_BYTES, + ).hex() + return SessionKeys(signing_key=SecretStr(signing)) + + +class NotSessionBearer(BaseModel): + """The bearer is not session-shaped; admission continues on its normal path.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["not_session_bearer"] = "not_session_bearer" + + +class SessionBearerAdmitted(BaseModel): + """A valid session access token: the principal to admit under after a live reload.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["admitted"] = "admitted" + principal: SessionPrincipal + + +class SessionBearerInvalid(BaseModel): + """The bearer is session-shaped but must not admit (expired, tampered, wrong key, or a + refresh token presented at the tool-call edge); admission fails closed with the + ``invalid_token`` challenge rather than falling through to another arm. ``expired`` + distinguishes a routine expiry (debug-log worthy) from a tampered or foreign token.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["invalid"] = "invalid" + expired: bool = False + + +SessionBearerResult: TypeAlias = NotSessionBearer | SessionBearerAdmitted | SessionBearerInvalid + + +def _strip_bearer(value: str) -> str: + parts = value.split(None, 1) + if len(parts) == 2 and parts[0].lower() == "bearer": + return parts[1] + return value + + +def is_session_bearer_shaped(authorization_value: str) -> bool: + """Cheap, keyless test that an ``Authorization`` value carries a session token of either + kind (optional ``Bearer`` scheme stripped). The admission edge engages the session arm + for an access token (to admit) and for a refresh token (to reject it explicitly, since + a refresh credential is never usable at the tool-call edge); anything else falls + through to normal admission.""" + candidate = _strip_bearer(authorization_value) + return is_session_token(candidate) or is_session_refresh_token(candidate) + + +def resolve_session_bearer( + authorization_value: str, + keys: SessionKeys, + now: datetime, +) -> SessionBearerResult: + """Classify an ``Authorization`` value presented at the aggregate MCP edge. + + Strips an optional ``Bearer`` scheme, then returns ``NotSessionBearer`` for a + non-session bearer (normal admission continues), ``SessionBearerAdmitted`` with the + recovered principal for a valid access token, and ``SessionBearerInvalid`` for a + session-shaped bearer that must not admit. Never raises: total over hostile input via + :func:`~.session_token.open_session_token`. + + A refresh token is ``SessionBearerInvalid`` here: it is a valid gateway credential but + only ever presented back to the token endpoint, so admission must fail it closed rather + than let it fall through to another arm. + """ + candidate = _strip_bearer(authorization_value) + if is_session_refresh_token(candidate): + return SessionBearerInvalid() + if not is_session_token(candidate): + return NotSessionBearer() + opened = open_session_token(candidate, keys, now) + if isinstance(opened, OpenedSessionToken): + return SessionBearerAdmitted(principal=opened.principal) + return SessionBearerInvalid(expired=isinstance(opened, SessionExpired)) + + +class SessionRefreshOpened(BaseModel): + """A valid session refresh token presented to the token endpoint: the principal to + re-validate and renew under.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["opened"] = "opened" + principal: SessionPrincipal + + +class SessionRefreshInvalid(BaseModel): + """The presented refresh grant is not a valid session refresh token for this client + (not refresh-shaped, will not open, or bound to a different ``client_id``); the token + endpoint fails the refresh closed.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["invalid"] = "invalid" + + +SessionRefreshResult: TypeAlias = SessionRefreshOpened | SessionRefreshInvalid + + +def open_session_refresh_bearer( + refresh_value: str, + keys: SessionKeys, + now: datetime, + expected_client_id: str, +) -> SessionRefreshResult: + """Open a session refresh token presented on a ``refresh_token`` grant. + + The token-endpoint mirror of :func:`resolve_session_bearer`: strips an optional + ``Bearer`` scheme, then returns ``SessionRefreshOpened`` with the recovered principal, + or ``SessionRefreshInvalid`` for anything that is not a valid session refresh token + issued to ``expected_client_id``. Never raises. The client binding (RFC 6749 section 6) + stops a refresh token stolen from one DCR client from being renewed through another; + ``client_id`` is not a secret (the caller presents it), so a plain equality check is + sufficient and, unlike ``hmac.compare_digest`` on ``str``, does not raise on non-ASCII. + """ + candidate = _strip_bearer(refresh_value) + if not is_session_refresh_token(candidate): + return SessionRefreshInvalid() + opened = open_session_refresh_token(candidate, keys, now) + if not isinstance(opened, OpenedSessionToken): + return SessionRefreshInvalid() + if opened.principal.client_id != expected_client_id: + return SessionRefreshInvalid() + return SessionRefreshOpened(principal=opened.principal) diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py new file mode 100644 index 00000000000..9325428f049 --- /dev/null +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/session_token.py @@ -0,0 +1,361 @@ +"""Identity-only session tokens for the gateway-level (aggregate ``/mcp``) DCR front door. + +A DCR client that signs in through LiteLLM SSO holds ONE bearer that carries ONLY a +litellm identity; unlike the :mod:`.envelope` bridge bearer it seals no upstream +credential, because the custody model vaults every upstream token server-side in +``LiteLLM_MCPUserCredentials`` and egress resolves them by user at call time. The token +is therefore a stable REFERENCE, not an authorization: admission reloads the live user +record and policy on every request, so deactivating the user (or their team) kills +outstanding sessions immediately without a revocation store. + +Wire shape: ``llm_session_`` (access) / ``llm_srefresh_`` (refresh) + an HS256 JWT, +the same signing approach as :mod:`.envelope`. Claims are ``iss``/``iat``/``exp`` +plus ``jti`` (per-mint uniqueness, so two tokens minted in the same second never +collide and a future revocation list has a stable handle), ``kind``, ``user_id``, and +``client_id``; ``client_id`` binds the refresh token +to the DCR client it was issued to (RFC 6749 section 6) and is carried on the access +token for parity and audit. There is no encrypted payload: nothing in a session token +is secret beyond the signature, and reprs never print the signed value because minted +tokens are ``SecretStr``. + +This module is pure and unwired: it imports nothing from endpoint or edge code, reads +no proxy globals, and takes all key material and the clock as explicit parameters. +Failures are values: :func:`open_session_token` and :func:`open_session_refresh_token` +are total over hostile, attacker-controlled input and return a +``SessionTokenOpenError`` variant rather than raising. PyJWT's ``iat``/``nbf``/``exp`` +validators are disabled for the same reasons documented in :mod:`.envelope` (they +raise on hostile claim types and compare against the wall clock instead of the +injected ``now``); the strict pydantic claims model is the sole, total type gate. +""" + +from __future__ import annotations + +import secrets +from datetime import datetime, timedelta +from typing import Literal, TypeAlias + +import jwt +from pydantic import BaseModel, ConfigDict, Field, SecretStr, ValidationError + +SESSION_TOKEN_PREFIX = "llm_session_" +"""Marker prefix on every serialized session ACCESS token so the admission edge can cheaply +tell a gateway session from a litellm key, JWT, or bridge envelope before doing any +cryptography. Distinct from the ``llm_env_``/``llm_refresh_`` envelope prefixes.""" + +SESSION_REFRESH_PREFIX = "llm_srefresh_" +"""Marker prefix on every serialized session REFRESH token. A distinct prefix keeps the two +credentials routable without crypto and, together with the signed ``kind`` claim, stops one +from being presented where the other is expected: the refresh token is only ever presented +back to the token endpoint, never at the MCP edge.""" + +SESSION_ISSUER = "litellm-mcp-gateway" +"""``iss`` claim stamped into every session token and required back on open. Distinct from +the envelope issuer so a token of one family can never validate in the other even under a +hypothetical shared signing key.""" + +SESSION_TTL_SECONDS = 3600 +"""Session ACCESS token lifetime (1h), matching the access-envelope and BYOK session bearer +windows: a client-held credential never outlives a bounded window, and each refresh +re-validates the live user before re-minting.""" + +SESSION_REFRESH_TTL_SECONDS = 1209600 +"""Session REFRESH token lifetime (14 days), matching the refresh-envelope bound. Each +renewal re-validates the sealed user against the live record (deactivation gates it) and +rotates the refresh token, so the practical bound is idle time, not a fixed session.""" + +MAX_SESSION_TOKEN_BYTES = 4096 +"""Size cap on the serialized token (prefix + JWT, in bytes) and on any candidate accepted +by the openers. Session claims are small; the only variable-length field is ``client_id`` +(a sealed DCR client record), and 4096 leaves ample headroom under common 8-16KB header +limits while bounding hostile input before JWT parsing.""" + +_SESSION_JWT_ALGORITHM = "HS256" + +SessionTokenKind = Literal["session", "session_refresh"] +"""Which credential a session token is. Stamped into the signed claims and required to match +on open, so a signature-valid token of one kind cannot be replayed as the other even if its +wire prefix is swapped (the prefix is not part of the signed payload; this claim is).""" + + +class SessionPrincipal(BaseModel): + """The litellm user a session token identifies and the DCR client it was issued to. + + ``user_id`` is the SSO-established litellm user subject, never a credential: admission + reloads the live user record by it, so current role, team, and revocation state are + enforced at use time rather than frozen at mint time. ``client_id`` is the (stateless, + gateway-sealed) DCR client identifier the token was issued to; the token endpoint + requires it to match on the refresh grant. + """ + + model_config = ConfigDict(frozen=True) + user_id: str = Field(min_length=1) + client_id: str = Field(min_length=1) + + +class SessionKeys(BaseModel): + """Injected key material: the HS256 signing key. + + ``signing_key`` must be at least 32 bytes: HS256's HMAC-SHA256 has a 256-bit security + level, RFC 7518 requires a key of at least that size, and a shorter key makes PyJWT + emit ``InsecureKeyLengthWarning``. + """ + + model_config = ConfigDict(frozen=True) + signing_key: SecretStr = Field(min_length=32) + + +class MintedSessionToken(BaseModel): + """A minted session token: the client-held bearer value and when it expires.""" + + model_config = ConfigDict(frozen=True) + token: SecretStr + expires_at: datetime + + +class OpenedSessionToken(BaseModel): + """A validated session token of either kind: the principal it was minted for.""" + + model_config = ConfigDict(frozen=True) + principal: SessionPrincipal + + +class SessionTokenTooLarge(BaseModel): + """The serialized token exceeded ``MAX_SESSION_TOKEN_BYTES``; carries sizes only. Only + reachable through an oversized ``client_id``, which registration should have bounded.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["session_token_too_large"] = "session_token_too_large" + size_bytes: int + max_bytes: int + + +SessionTokenMintError: TypeAlias = SessionTokenTooLarge + + +class NotASessionToken(BaseModel): + """The candidate does not carry the expected session prefix.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["not_a_session_token"] = "not_a_session_token" + + +class SessionBadSignature(BaseModel): + """The JWT signature does not verify under the provided signing key.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["session_bad_signature"] = "session_bad_signature" + + +class SessionExpired(BaseModel): + """The token's ``exp`` is not in the future relative to the provided ``now``.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["session_expired"] = "session_expired" + + +class SessionMalformed(BaseModel): + """The token is not a well-formed session token: undecodable JWT, wrong issuer, wrong + ``kind``, or missing/mistyped/extra claims.""" + + model_config = ConfigDict(frozen=True) + tag: Literal["session_malformed"] = "session_malformed" + + +SessionTokenOpenError: TypeAlias = NotASessionToken | SessionBadSignature | SessionExpired | SessionMalformed + + +class _SessionClaims(BaseModel): + """Decoded-claims boundary that pins the exact shape the mints emit. + + ``user_id``/``client_id`` mirror the ``min_length`` constraints of + :class:`SessionPrincipal` so any claim set that validates here also constructs a + principal, keeping the openers raise-free: a correctly signed JWT with an empty + identity claim fails here and maps to ``SessionMalformed``. ``strict`` rejects coerced + types (``exp: "123"``) and ``extra="forbid"`` rejects any claim the gateway never + mints; PyJWT's own registered-claim validators are disabled at decode (see module + docstring), so this model is the sole, total type gate for every claim. + """ + + model_config = ConfigDict(frozen=True, strict=True, extra="forbid") + iss: str + iat: int + exp: int + jti: str = Field(min_length=1) + kind: SessionTokenKind + user_id: str = Field(min_length=1) + client_id: str = Field(min_length=1) + + +def is_session_token(candidate: str) -> bool: + """Cheap prefix check for a session ACCESS token so the admission edge can route gateway + sessions vs keys, JWTs, and envelopes without crypto.""" + return candidate.startswith(SESSION_TOKEN_PREFIX) + + +def is_session_refresh_token(candidate: str) -> bool: + """Cheap prefix check for a session REFRESH token so the token endpoint can route a + refresh grant without crypto.""" + return candidate.startswith(SESSION_REFRESH_PREFIX) + + +def mint_session_token( + principal: SessionPrincipal, + keys: SessionKeys, + now: datetime, +) -> MintedSessionToken | SessionTokenMintError: + """Mint the short-lived session ACCESS token for ``principal``. + + ``exp`` is ``SESSION_TTL_SECONDS`` from ``now``. Returns ``SessionTokenTooLarge`` when + the serialized token exceeds ``MAX_SESSION_TOKEN_BYTES``. + """ + return _mint( + kind="session", + prefix=SESSION_TOKEN_PREFIX, + principal=principal, + expires_at=now + timedelta(seconds=SESSION_TTL_SECONDS), + keys=keys, + now=now, + ) + + +def mint_session_refresh_token( + principal: SessionPrincipal, + keys: SessionKeys, + now: datetime, +) -> MintedSessionToken | SessionTokenMintError: + """Mint the long-lived session REFRESH token for ``principal``. + + ``exp`` is ``SESSION_REFRESH_TTL_SECONDS`` from ``now``. Minting a distinct + ``kind="session_refresh"`` claim is what keeps a refresh token from ever opening as an + access credential at the MCP edge. + """ + return _mint( + kind="session_refresh", + prefix=SESSION_REFRESH_PREFIX, + principal=principal, + expires_at=now + timedelta(seconds=SESSION_REFRESH_TTL_SECONDS), + keys=keys, + now=now, + ) + + +def open_session_token( + candidate: str, + keys: SessionKeys, + now: datetime, +) -> OpenedSessionToken | SessionTokenOpenError: + """Validate a session ACCESS ``candidate`` and recover the principal. + + Never raises for bad input: every invalid, expired, tampered, or wrong-kind candidate + maps to a distinct ``SessionTokenOpenError`` variant. + """ + return _open(candidate, prefix=SESSION_TOKEN_PREFIX, expected_kind="session", keys=keys, now=now) + + +def open_session_refresh_token( + candidate: str, + keys: SessionKeys, + now: datetime, +) -> OpenedSessionToken | SessionTokenOpenError: + """Validate a session REFRESH ``candidate`` and recover the principal. + + Total over hostile input exactly like :func:`open_session_token`. The + ``kind="session_refresh"`` claim is required, so an access token re-prefixed as a + refresh one is rejected as ``SessionMalformed``. + """ + return _open(candidate, prefix=SESSION_REFRESH_PREFIX, expected_kind="session_refresh", keys=keys, now=now) + + +def _mint( + kind: SessionTokenKind, + prefix: str, + principal: SessionPrincipal, + expires_at: datetime, + keys: SessionKeys, + now: datetime, +) -> MintedSessionToken | SessionTokenTooLarge: + """Sign the claims for either token kind and enforce the size cap. Shared by both mints + so the JWT shape, issuer, and size guard cannot drift between access and refresh.""" + claims = _SessionClaims( + iss=SESSION_ISSUER, + iat=int(now.timestamp()), + exp=int(expires_at.timestamp()), + jti=secrets.token_urlsafe(16), + kind=kind, + user_id=principal.user_id, + client_id=principal.client_id, + ) + token = prefix + jwt.encode( + claims.model_dump(), keys.signing_key.get_secret_value(), algorithm=_SESSION_JWT_ALGORITHM + ) + size_bytes = len(token.encode("utf-8")) + if size_bytes > MAX_SESSION_TOKEN_BYTES: + return SessionTokenTooLarge(size_bytes=size_bytes, max_bytes=MAX_SESSION_TOKEN_BYTES) + return MintedSessionToken(token=SecretStr(token), expires_at=expires_at) + + +def _open( + candidate: str, + prefix: str, + expected_kind: SessionTokenKind, + keys: SessionKeys, + now: datetime, +) -> OpenedSessionToken | SessionTokenOpenError: + """Prefix-route, size-bound, signature-verify, kind-check, and expiry-check an + attacker-controlled candidate, shared by both openers so the security gate is identical + for access and refresh. Returns the opened token or a distinct error; never raises.""" + if not candidate.startswith(prefix): + return NotASessionToken() + # UTF-8 byte length is never below character length, so a character count already over + # the cap rejects an oversize candidate in O(1) without encoding it; the exact byte + # check then runs only on candidates already bounded to the cap in characters. + if len(candidate) > MAX_SESSION_TOKEN_BYTES: + return SessionMalformed() + if len(candidate.encode("utf-8", "surrogatepass")) > MAX_SESSION_TOKEN_BYTES: + return SessionMalformed() + claims = _decode_claims(candidate.removeprefix(prefix), keys.signing_key) + if not isinstance(claims, _SessionClaims): + return claims + if claims.kind != expected_kind: + return SessionMalformed() + if now.timestamp() >= claims.exp: + return SessionExpired() + return OpenedSessionToken(principal=SessionPrincipal(user_id=claims.user_id, client_id=claims.client_id)) + + +def _decode_claims( + compact: str, + signing_key: SecretStr, +) -> _SessionClaims | SessionBadSignature | SessionMalformed: + """Verify the HS256 signature and shape of an attacker-controlled compact JWT. + + ``compact`` is fully hostile and bounded to ``MAX_SESSION_TOKEN_BYTES`` by the caller. + PyJWT's ``iat``/``nbf``/``exp`` validators are disabled: they raise on hostile claim + types and, for ``iat``/``nbf``, compare against the wall clock rather than the injected + ``now`` (``exp`` is checked by the caller against ``now``). Apart from a signature + mismatch, every decode failure is ``SessionMalformed``: a non-UTF-8 candidate surfaces + as ``UnicodeEncodeError`` (a ``ValueError``), a non-string registered claim as a + ``TypeError`` from PyJWT's claim validators, and a wrong issuer or structurally invalid + token as an ``InvalidTokenError``. ``_SessionClaims`` is the total type gate. + """ + try: + payload = jwt.decode( + compact, + signing_key.get_secret_value(), + algorithms=[_SESSION_JWT_ALGORITHM], + issuer=SESSION_ISSUER, + options={ + "verify_exp": False, + "verify_iat": False, + "verify_nbf": False, + "require": ["iss", "iat", "exp"], + }, + ) + except jwt.InvalidSignatureError: + return SessionBadSignature() + except (jwt.InvalidTokenError, ValueError, TypeError): + return SessionMalformed() + try: + return _SessionClaims.model_validate(payload) + except ValidationError: + return SessionMalformed() diff --git a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py index 64a20255ab2..926d96c8868 100644 --- a/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py +++ b/litellm/proxy/_experimental/mcp_server/outbound_credentials/types.py @@ -184,7 +184,12 @@ class ClientCredentialsConfig(BaseModel): Fields are optional so the config can be built incomplete: a value may be supplied at runtime (`token_url` via RFC 8414 discovery, `client_id`/`secret` via DCR), and the - resolver arm raises `CredError.misconfigured` when a needed field is still absent. + resolver arm returns `CredError.misconfigured` when a needed field is still absent. + + `audience` is the IdP-specific audience parameter some authorization servers require on + the client_credentials grant (sent as `audience` in the token request when set). + `token_endpoint_auth_method` selects how the client authenticates to the token endpoint + (RFC 6749 section 2.3.1); `None` defaults to `client_secret_post`. """ model_config = ConfigDict(frozen=True) @@ -193,6 +198,8 @@ class ClientCredentialsConfig(BaseModel): client_secret: SecretStr | None = None token_url: str | None = None scopes: tuple[str, ...] = () + audience: str | None = None + token_endpoint_auth_method: Literal["client_secret_post", "client_secret_basic"] | None = None class TokenExchangeConfig(BaseModel): diff --git a/litellm/proxy/auth/auth_utils.py b/litellm/proxy/auth/auth_utils.py index 38900260c98..293bb74e211 100644 --- a/litellm/proxy/auth/auth_utils.py +++ b/litellm/proxy/auth/auth_utils.py @@ -273,6 +273,7 @@ _BANNED_REQUEST_BODY_PARAMS: Tuple[str, ...] = ( # re-route the request's retention and accounting to any project # reachable with the deployment's shared AWS credentials. "aws_bedrock_project_id", + "bedrock_tags", # Provider-specific endpoint overrides that flow into the outbound # request via ``optional_params``. Same threat as ``api_base``: # ``s3_endpoint_url`` redirects Bedrock file uploads to attacker diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 5a0b94d1524..55b50e7d9ff 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -1647,7 +1647,7 @@ async def ui_view_spend_logs( description="Time till which to view key spend", ), page: int = fastapi.Query(default=1, description="Page number for pagination", ge=1), - page_size: int = fastapi.Query(default=50, description="Number of items per page", ge=1, le=100), + page_size: int = fastapi.Query(default=50, description="Number of items per page", ge=1, le=1000), user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth), status_filter: str | None = fastapi.Query( default=None, description="Filter logs by status (e.g., success, failure)" diff --git a/litellm/responses/mcp/chat_completions_handler.py b/litellm/responses/mcp/chat_completions_handler.py index f2ccfd430ae..5c3e0cf0902 100644 --- a/litellm/responses/mcp/chat_completions_handler.py +++ b/litellm/responses/mcp/chat_completions_handler.py @@ -12,7 +12,7 @@ from typing import ( from litellm.responses.mcp.litellm_proxy_mcp_handler import ( LiteLLM_Proxy_MCP_Handler, ) -from litellm.responses.utils import ResponsesAPIRequestUtils +from litellm.responses.mcp.request_context import MCPRequestContext from litellm.types.utils import ModelResponse from litellm.utils import CustomStreamWrapper @@ -114,20 +114,13 @@ async def acompletion_with_mcp( **kwargs, ) - # Extract user_api_key_auth from metadata or kwargs - user_api_key_auth = kwargs.get("user_api_key_auth") or ((kwargs.get("metadata", {}) or {}).get("user_api_key_auth")) - request_tags = LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(kwargs) - - # Extract MCP auth headers before fetching tools (needed for dynamic auth) - ( - mcp_auth_header, - mcp_server_auth_headers, - oauth2_headers, - raw_headers, - ) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request( - secret_fields=kwargs.get("secret_fields"), - tools=tools, - ) + context = MCPRequestContext.resolve(kwargs=kwargs, tools=tools) + user_api_key_auth = context.user_api_key_auth + request_tags = list(context.request_tags) if context.request_tags else None + mcp_auth_header = context.mcp_auth_header + mcp_server_auth_headers = context.mcp_server_auth_headers + oauth2_headers = context.oauth2_headers + raw_headers = context.raw_headers # Process MCP tools (pass auth headers for dynamic auth) ( diff --git a/litellm/responses/mcp/litellm_proxy_mcp_handler.py b/litellm/responses/mcp/litellm_proxy_mcp_handler.py index 392bb7bcab2..cb680dc8b86 100644 --- a/litellm/responses/mcp/litellm_proxy_mcp_handler.py +++ b/litellm/responses/mcp/litellm_proxy_mcp_handler.py @@ -479,17 +479,25 @@ class LiteLLM_Proxy_MCP_Handler: ) -> bool: """Check if we should auto-execute tool calls. - Only auto-execute tools if user passed a MCP tool with require_approval set to "never". - - + Auto-execution requires EVERY MCP reference to opt in with + ``require_approval="never"``. A single reference that requires approval + ("always", "manual", the object form, or an unset value) disables + auto-execution for the whole request. This fails closed: when an + approval-required reference shares a request with a "never" one, the + model's tool calls are returned to the caller instead of being run, so + an approval-gated tool can never be invoked without approval. Returns + False for an empty list. """ - for tool in mcp_tools_with_litellm_proxy: - if isinstance(tool, dict): - if tool.get("require_approval") == "never": - return True - elif getattr(tool, "require_approval", None) == "never": - return True - return False + references = list(mcp_tools_with_litellm_proxy or []) + if not references: + return False + for tool in references: + approval = ( + tool.get("require_approval") if isinstance(tool, dict) else getattr(tool, "require_approval", None) + ) + if approval != "never": + return False + return True @staticmethod def _extract_tool_calls_from_response(response: ResponsesAPIResponse) -> List[Any]: @@ -542,7 +550,10 @@ class LiteLLM_Proxy_MCP_Handler: tool_arguments = function_block.get("arguments") else: tool_name = tool_call.get("name") + # Anthropic tool_use blocks carry the arguments under `input` tool_arguments = tool_call.get("arguments") + if tool_arguments is None: + tool_arguments = tool_call.get("input") else: tool_call_id = getattr(tool_call, "call_id", None) or getattr(tool_call, "id", None) @@ -553,6 +564,8 @@ class LiteLLM_Proxy_MCP_Handler: else: tool_name = getattr(tool_call, "name", None) tool_arguments = getattr(tool_call, "arguments", None) + if tool_arguments is None: + tool_arguments = getattr(tool_call, "input", None) return tool_name, tool_arguments, tool_call_id diff --git a/litellm/responses/mcp/request_context.py b/litellm/responses/mcp/request_context.py new file mode 100644 index 00000000000..fa03e677b39 --- /dev/null +++ b/litellm/responses/mcp/request_context.py @@ -0,0 +1,73 @@ +""" +The per-request context an MCP gateway handler needs. + +Listing and executing MCP tools both need the caller's identity, their MCP auth +headers, and the request's trace/tag identifiers. Every gateway surface resolves +the same set from its own kwargs, so resolving it in one place keeps a new +surface from silently dropping a field: omitting the auth headers, for instance, +still executes the tool, just with no credentials. +""" + +from dataclasses import dataclass +from typing import Any, Iterable, Mapping, Sequence, Union + + +@dataclass(frozen=True, slots=True) +class MCPRequestContext: + """Everything a gateway handler must forward to MCP tool listing and execution.""" + + user_api_key_auth: Any # any-ok: UserAPIKeyAuth is proxy-only; importing it here would create a cycle + mcp_auth_header: Union[str, None] = None + mcp_server_auth_headers: Union[Mapping[str, Mapping[str, str]], None] = None + oauth2_headers: Union[Mapping[str, str], None] = None + raw_headers: Union[Mapping[str, str], None] = None + request_tags: Union[Sequence[str], None] = None + litellm_trace_id: Union[str, None] = None + litellm_call_id: Union[str, None] = None + + @classmethod + def resolve( + cls, + kwargs: Mapping[str, Any], + tools: Union[Iterable[Any], None], + ) -> "MCPRequestContext": + """ + Build the context from a gateway handler's kwargs. + + ``user_api_key_auth`` is read from both metadata keys because routes differ: + LITELLM_METADATA_ROUTES (``/v1/messages``, ``/responses``) carry it in + ``litellm_metadata`` while ``/chat/completions`` uses ``metadata``. + """ + from litellm.responses.mcp.litellm_proxy_mcp_handler import ( + LiteLLM_Proxy_MCP_Handler, + ) + from litellm.responses.utils import ResponsesAPIRequestUtils + + litellm_metadata = kwargs.get("litellm_metadata") or {} + metadata = kwargs.get("metadata") or {} + user_api_key_auth = ( + kwargs.get("user_api_key_auth") + or litellm_metadata.get("user_api_key_auth") + or metadata.get("user_api_key_auth") + ) + + ( + mcp_auth_header, + mcp_server_auth_headers, + oauth2_headers, + raw_headers, + ) = ResponsesAPIRequestUtils.extract_mcp_headers_from_request( + secret_fields=kwargs.get("secret_fields"), + tools=tools, + ) + + return cls( + user_api_key_auth=user_api_key_auth, + mcp_auth_header=mcp_auth_header, + mcp_server_auth_headers=mcp_server_auth_headers, + oauth2_headers=oauth2_headers, + raw_headers=raw_headers, + request_tags=LiteLLM_Proxy_MCP_Handler._get_parent_request_tags(dict(kwargs)), + litellm_trace_id=kwargs.get("litellm_trace_id"), + litellm_call_id=kwargs.get("litellm_call_id"), + ) diff --git a/litellm/rust_bridge/ocr.py b/litellm/rust_bridge/ocr.py index e9139a634f1..91aac6c1232 100644 --- a/litellm/rust_bridge/ocr.py +++ b/litellm/rust_bridge/ocr.py @@ -72,17 +72,28 @@ def use_litellm_rust( messages: RustMessages | None | _Unset = _UNSET, amessages: RustAmessages | None | _Unset = _UNSET, responses_websocket: Any | None | _Unset = _UNSET, + transcription: Any | None | _Unset = _UNSET, + atranscription: Any | None | _Unset = _UNSET, ) -> None: global _rust_ocr_enabled, _rust_ocr_impl, _rust_aocr_impl configuring_ocr = not isinstance(ocr, _Unset) or not isinstance(aocr, _Unset) configuring_messages = not isinstance(messages, _Unset) or not isinstance(amessages, _Unset) configuring_responses_websocket = not isinstance(responses_websocket, _Unset) + configuring_transcription = not isinstance(transcription, _Unset) or not isinstance(atranscription, _Unset) if configuring_ocr or (not configuring_messages and not configuring_responses_websocket): _rust_ocr_enabled = enabled if not isinstance(ocr, _Unset): _rust_ocr_impl = ocr if not isinstance(aocr, _Unset): _rust_aocr_impl = aocr + if configuring_transcription: + from litellm.rust_bridge.transcription import configure_rust_transcription + + configure_rust_transcription( + enabled=enabled, + transcription=transcription, + atranscription=atranscription, + ) if not configuring_messages and not configuring_responses_websocket: return if configuring_messages: diff --git a/litellm/rust_bridge/transcription.py b/litellm/rust_bridge/transcription.py new file mode 100644 index 00000000000..44bb42e4104 --- /dev/null +++ b/litellm/rust_bridge/transcription.py @@ -0,0 +1,148 @@ +from __future__ import annotations + +from dataclasses import dataclass +from typing import Awaitable, Final, Protocol, Union, cast + +import httpx + +from litellm.rust_bridge.timeouts import timeout_to_seconds + + +class RustTranscription(Protocol): + def __call__( + self, + model: str, + audio: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: + raise NotImplementedError + + +class RustAtranscription(Protocol): + def __call__( + self, + model: str, + audio: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> Awaitable[dict[str, object]]: + raise NotImplementedError + + +class _Unset: + pass + + +_UNSET: Final[_Unset] = _Unset() + + +@dataclass +class _RustTranscriptionState: + transcription: RustTranscription | None = None + atranscription: RustAtranscription | None = None + + +_STATE = _RustTranscriptionState() + + +def configure_rust_transcription( + enabled: bool = True, + *, + transcription: RustTranscription | None | _Unset = _UNSET, + atranscription: RustAtranscription | None | _Unset = _UNSET, +) -> None: + if not isinstance(transcription, _Unset): + _STATE.transcription = transcription + if not isinstance(atranscription, _Unset): + _STATE.atranscription = atranscription + + +def load_rust_transcription() -> RustTranscription | None: + if _STATE.transcription is not None: + return _STATE.transcription + from litellm.rust_bridge import get_native_bridge + + native_bridge = get_native_bridge() + return ( + None + if native_bridge is None + else cast( # cast-ok: native extension protocol is runtime-defined + RustTranscription, getattr(native_bridge, "transcription", None) + ) + ) + + +def load_rust_atranscription() -> RustAtranscription | None: + if _STATE.atranscription is not None: + return _STATE.atranscription + from litellm.rust_bridge import get_native_bridge + + native_bridge = get_native_bridge() + return ( + None + if native_bridge is None + else cast( # cast-ok: native extension protocol is runtime-defined + RustAtranscription, getattr(native_bridge, "atranscription", None) + ) + ) + + +def transcription( + *, + model: str, + audio: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout: Union[float, httpx.Timeout] | None, +) -> dict[str, object] | None: + rust_transcription = load_rust_transcription() + if rust_transcription is None: + return None + return rust_transcription( + model=model, + audio=audio, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout_seconds=timeout_to_seconds(timeout), + ) + + +async def atranscription( + *, + model: str, + audio: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout: Union[float, httpx.Timeout] | None, +) -> dict[str, object] | None: + rust_atranscription = load_rust_atranscription() + if rust_atranscription is None: + return None + return await rust_atranscription( + model=model, + audio=audio, + api_key=api_key, + api_base=api_base, + custom_llm_provider=custom_llm_provider, + extra_headers=extra_headers, + optional_params=optional_params, + timeout_seconds=timeout_to_seconds(timeout), + ) diff --git a/litellm/types/integrations/anthropic_cache_control_hook.py b/litellm/types/integrations/anthropic_cache_control_hook.py index 601978bb04f..efb189088b6 100644 --- a/litellm/types/integrations/anthropic_cache_control_hook.py +++ b/litellm/types/integrations/anthropic_cache_control_hook.py @@ -1,6 +1,6 @@ from typing import Literal, Optional, Union -from typing_extensions import TypedDict +from typing_extensions import NotRequired, TypedDict from litellm.types.llms.openai import ChatCompletionCachedContent @@ -12,6 +12,7 @@ class CacheControlMessageInjectionPoint(TypedDict): role: Optional[Literal["user", "system", "assistant"]] # Optional: target by role (user, system, assistant) index: Optional[Union[int, str]] # Optional: target by specific index control: Optional[ChatCompletionCachedContent] + _litellm_judged: NotRequired[bool] # Internal: written back by litellm once the client cache_control judgment ran class CacheControlToolConfigInjectionPoint(TypedDict): @@ -19,6 +20,7 @@ class CacheControlToolConfigInjectionPoint(TypedDict): location: Literal["tool_config"] control: Optional[ChatCompletionCachedContent] + _litellm_judged: NotRequired[bool] # Internal: written back by litellm once the client cache_control judgment ran CacheControlInjectionPoint = Union[ diff --git a/litellm/types/llms/bedrock.py b/litellm/types/llms/bedrock.py index bdf6b8fefed..d9f8229dbed 100644 --- a/litellm/types/llms/bedrock.py +++ b/litellm/types/llms/bedrock.py @@ -985,6 +985,11 @@ class BedrockOutputDataConfig(TypedDict): s3OutputDataConfig: BedrockS3OutputDataConfig +class BedrockTag(TypedDict): + key: str + value: str + + class BedrockCreateBatchRequest(TypedDict, total=False): """ Request structure for creating a Bedrock batch inference job. @@ -999,7 +1004,7 @@ class BedrockCreateBatchRequest(TypedDict, total=False): outputDataConfig: BedrockOutputDataConfig timeoutDurationInHours: Optional[int] clientRequestToken: Optional[str] - tags: Optional[List[dict]] + tags: Optional[List[BedrockTag]] BedrockBatchJobStatus = Literal["Submitted", "InProgress", "Completed", "Failed", "Stopping", "Stopped"] diff --git a/tests/e2e/access_control/test_access_control_e2e.py b/tests/e2e/access_control/test_access_control_e2e.py index ce649fa2400..e24b721d831 100644 --- a/tests/e2e/access_control/test_access_control_e2e.py +++ b/tests/e2e/access_control/test_access_control_e2e.py @@ -23,12 +23,16 @@ from access_control_client import ( ROUTE_NOT_ALLOWED_MARKER, ) from e2e_config import unique_marker +from e2e_http import Success, UnauthorizedError, UnknownApiError, unwrap from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, LiteLLMParamsBody +from proxy_client import ProxyClient pytestmark = pytest.mark.e2e ALLOWED_MODEL = "gemini-2.5-flash" DISALLOWED_MODEL = "gpt-5.5" +VIRTUAL_KEY_BACKEND = "anthropic/claude-haiku-4-5-20251001" def _is_json(body: str) -> bool: @@ -39,6 +43,7 @@ def _is_json(body: str) -> bool: return False + class TestAccessControl: def test_disallowed_model_is_denied_403( self, client: AccessControlClient, resources: ResourceManager @@ -81,3 +86,59 @@ class TestAccessControl: f"{result.status_code}: {result.body[:300]}" ) assert _is_json(result.body), f"400 body must be valid JSON: {result.body[:300]}" + + +class TestVirtualKeyAuth: + """Virtual-key auth the way OpenAI-compatible clients send it: a real key + must reach chat, a forged bearer must be rejected before the provider.""" + + @pytest.mark.covers( + "mgmt.virtual_key.valid_allows", + "mgmt.virtual_key.invalid_denied", + exercised_on=[], + ) + def test_valid_key_allows_and_invalid_key_denied( + self, proxy: ProxyClient, resources: ResourceManager + ) -> None: + model = f"e2e-auth-chat-{unique_marker()}" + model_id = proxy.create_model( + model, + LiteLLMParamsBody(model=VIRTUAL_KEY_BACKEND, api_key="os.environ/ANTHROPIC_API_KEY"), + ) + resources.defer(lambda: proxy.delete_model(model_id)) + key = resources.key() + + ok = unwrap( + proxy.chat( + key, + ChatBody( + model=model, + messages=[ + ChatMessage( + role="user", + content=f"Reply with one word. {unique_marker()}", + ) + ], + max_tokens=16, + ), + ) + ) + assert ok.choices, f"valid key must complete chat: {ok}" + + bad = proxy.chat( + "sk-e2e-forged-not-a-real-key", + ChatBody( + model=model, + messages=[ChatMessage(role="user", content="should not run")], + max_tokens=8, + ), + ) + match bad: + case UnauthorizedError(): + return + case UnknownApiError(status_code=status) if status in (401, 403): + return + case Success(): + pytest.fail("forged bearer must not reach a successful completion") + case _: + pytest.fail(f"forged bearer must be auth-denied, got {bad}") diff --git a/tests/e2e/batches/test_batches_e2e.py b/tests/e2e/batches/test_batches_e2e.py index 8f10c8c7c2a..f886c09b705 100644 --- a/tests/e2e/batches/test_batches_e2e.py +++ b/tests/e2e/batches/test_batches_e2e.py @@ -16,13 +16,14 @@ misroute to the wrong provider fails the create. from __future__ import annotations import json +import os import time from datetime import datetime, timedelta, timezone from typing import Callable import pytest -from e2e_config import unique_marker +from e2e_config import require_env, unique_marker from batch_client import ( BatchClient, @@ -39,7 +40,9 @@ from capabilities import ( FILE_ID_SHAPE, OPENAI_BATCH_MODEL, Capability, + batch_model_name, coverage_cells_for_lifecycle, + is_managed_id, matches_id_shape, raw_id_matches_provider, ) @@ -53,7 +56,7 @@ from e2e_http import ( unwrap, ) from lifecycle import ResourceManager -from models import KeyGenerateBody, SpendLogRow +from models import KeyGenerateBody, LiteLLMParamsBody, SpendLogRow pytestmark = pytest.mark.e2e @@ -457,3 +460,328 @@ def test_rate_limited_batch_create_leaves_no_unattributed_spend_row( "batch create on a rate-limited key left an unattributed spend row " f"(LIT-3266); rows={[(r.request_id, r.call_type, r.model) for r in new_orphans]}" ) + + +OPENAI_FILE_CONTENT_BACKEND = "gpt-4o-mini" + + +class TestBatchFileContent: + """GET /v1/files/{id}/content returns the uploaded batch JSONL bytes.""" + + @pytest.mark.covers( + "llm.files.openai.content.nonstream.works", + exercised_on=["files"], + ) + def test_file_content_matches_upload( + self, client: BatchClient, resources: ResourceManager + ) -> None: + proxy_name = f"e2e-file-content-{unique_marker()}" + model_id = client.create_model( + proxy_name, + LiteLLMParamsBody( + model=f"openai/{OPENAI_FILE_CONTENT_BACKEND}", + api_key="os.environ/OPENAI_API_KEY", + ), + ) + resources.defer(lambda: client.delete_model(model_id)) + key = resources.key() + + payload = render_jsonl(OPENAI_FILE_CONTENT_BACKEND) + file = unwrap( + client.upload_file( + content=payload, + form=FileUploadForm(purpose="batch", target_model_names=proxy_name), + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + assert file.id + + downloaded = client.proxy.transport.download( + f"/v1/files/{file.id}/content", + headers=client.proxy.transport.bearer(key), + ) + assert downloaded.status_code == 200, ( + f"file content must be 200, got {downloaded.status_code}: {downloaded.body[:300]}" + ) + expected = payload.decode().rstrip("\n") + got = downloaded.body.rstrip("\n") + assert got == expected, ( + "downloaded file content must match the uploaded JSONL bytes" + ) + + +BATCH_RL_REQUEST_LINES = 3 +BATCH_RL_RPM_LIMIT = 2 + + +def _multi_request_jsonl(model: str, n: int) -> bytes: + lines = tuple( + json.dumps( + { + "custom_id": f"req-{i}", + "method": "POST", + "url": "/v1/chat/completions", + "body": { + "model": model, + "messages": [{"role": "user", "content": "ping"}], + "max_tokens": 8, + }, + } + ) + for i in range(n) + ) + return ("\n".join(lines) + "\n").encode() + + +class TestBatchRateLimitErrorMapping: + """Batch create that exceeds a key's RPM maps to a structured 429. + + The batch rate limiter reads the input file at submission time and rejects + the create when the file's request count would exceed the key's remaining + RPM. The product promise is not only the block itself but the + OpenAI-compatible shape: HTTP 429, a body that names the batch rate limit, + and pacing headers so clients can back off. Complements the LIT-3266 hygiene + check (no orphan spend rows) by asserting the error mapping when the limiter + actually fires. + """ + + @pytest.mark.covers( + "quota_management.ratelimit.batch_rpm.blocks_over_limit", + exercised_on=["batches"], + ) + def test_batch_create_over_rpm_returns_mapped_429( + self, client: BatchClient, resources: ResourceManager, batch_deployments: None + ) -> None: + user_id = f"e2e-batch-rl-map-{unique_marker()}" + key = client.proxy.generate_key( + KeyGenerateBody( + models=[], rpm_limit=BATCH_RL_RPM_LIMIT, tpm_limit=1_000_000, user_id=user_id + ) + ) + resources.defer(lambda: client.proxy.delete_key(key)) + + file = unwrap( + client.upload_file( + content=_multi_request_jsonl("gpt-4o-mini", BATCH_RL_REQUEST_LINES), + form=FileUploadForm(purpose="batch"), + model=OPENAI_BATCH_MODEL, + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + + created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + + assert created.status_code == 429, ( + f"expected batch RPM 429 when file has {BATCH_RL_REQUEST_LINES} requests and " + f"rpm_limit={BATCH_RL_RPM_LIMIT}, got {created.status_code}: {created.body[:400]}" + ) + body_lower = created.body.lower() + assert "batch rate limit exceeded" in body_lower, ( + f"429 body must name the batch rate limit so clients can branch on it; " + f"got: {created.body[:400]}" + ) + assert str(BATCH_RL_REQUEST_LINES) in created.body, ( + f"429 body should report the batch request count ({BATCH_RL_REQUEST_LINES}); " + f"got: {created.body[:400]}" + ) + assert "rpm" in body_lower or "requests remaining" in body_lower, ( + f"429 body must describe the RPM budget remaining so clients can pace; " + f"got: {created.body[:400]}" + ) + retry_after = created.headers.get("retry-after") + if retry_after is not None: + assert retry_after.isdigit() and int(retry_after) > 0, ( + f"retry-after must be a positive integer when present, got {retry_after!r}" + ) + + +ASSUME_ROLE_RAW_MODEL = "bedrock/us.anthropic.claude-haiku-4-5-20251001-v1:0" + + +def _assume_role_params(role_arn: str, session_name: str) -> LiteLLMParamsBody: + return LiteLLMParamsBody( + model=ASSUME_ROLE_RAW_MODEL, + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + s3_region_name="os.environ/AWS_REGION", + s3_bucket_name="os.environ/AWS_BATCH_S3_BUCKET", + s3_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + s3_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_batch_role_arn="os.environ/AWS_BATCH_ROLE_ARN", + aws_role_name=role_arn, + aws_session_name=session_name, + ) + + +class TestBedrockBatchAssumeRole: + """Bedrock batch create under STS assume-role credentials. + + Provisions a bedrock batch deployment whose litellm_params carry + aws_role_name / aws_session_name (the product path for role assumption) and + runs the unified file-upload + batch-create lifecycle. Success means the + proxy assumed the role and Bedrock accepted the job; a misconfigured role + fails create with an AWS auth error rather than silently falling back to the + ambient key. + """ + + @pytest.mark.covers( + "llm.batches.bedrock.assume_role.nonstream.works", + "llm.files.bedrock.upload.nonstream.works", + exercised_on=["batches", "files"], + ) + def test_unified_batch_create_with_assume_role( + self, client: BatchClient, resources: ResourceManager + ) -> None: + (role_arn,) = require_env("AWS_ROLE_NAME") + require_env( + "AWS_ACCESS_KEY_ID", + "AWS_SECRET_ACCESS_KEY", + "AWS_REGION", + "AWS_BATCH_S3_BUCKET", + "AWS_BATCH_ROLE_ARN", + ) + session_name = f"e2e-batch-sts-{unique_marker()}"[:64] + model_name = batch_model_name("bedrock-sts-batch") + + model_id = client.create_model(model_name, _assume_role_params(role_arn, session_name)) + resources.defer(lambda: client.delete_model(model_id)) + key = resources.key() + + file = unwrap( + client.upload_file( + content=render_jsonl(ASSUME_ROLE_RAW_MODEL), + form=FileUploadForm(purpose="batch", target_model_names=model_name), + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + assert_file_object(file, provider="bedrock") + + created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + require_successful_call(created) + batch = BatchObject.model_validate_json(created.body) + resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key))) + + assert batch.id, f"assume-role create returned no batch id: {created.body[:200]}" + assert is_managed_id(batch.id), ( + f"assume-role create via target_model_names must return a managed batch id, " + f"got {batch.id!r}" + ) + assert batch.status in CREATED_BATCH_STATUSES, ( + f"assume-role batch has non-transitional status {batch.status!r}" + ) + assert_batch_object(batch) + + fetched = unwrap(client.retrieve_batch(batch.id, key=key)) + assert fetched.id == batch.id + + +GEMINI_FILES_RAW_MODEL = "gemini-2.5-flash" + + +class TestGeminiFiles: + """Gemini Files API upload through the proxy (LIT-3382). + + gemini is a first-class FileCreateProvider. The test registers a gemini + deployment, uploads a tiny batch-purpose JSONL with target_model_names + routing, and asserts a FileObject comes back. Batch create for pure gemini + (non-Vertex) is out of scope here; Vertex covers the Gemini batch job path in + the main lifecycle matrix. + """ + + @pytest.mark.covers( + "llm.files.gemini.upload.nonstream.works", + exercised_on=["files"], + ) + def test_gemini_file_upload( + self, client: BatchClient, resources: ResourceManager + ) -> None: + model_name = batch_model_name("gemini-files") + model_id = client.create_model( + model_name, + LiteLLMParamsBody( + model=f"gemini/{GEMINI_FILES_RAW_MODEL}", + api_key="os.environ/GEMINI_API_KEY", + ), + ) + resources.defer(lambda: client.delete_model(model_id)) + key = resources.key() + + file = unwrap( + client.upload_file( + content=render_jsonl(GEMINI_FILES_RAW_MODEL), + form=FileUploadForm(purpose="batch", target_model_names=model_name), + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + assert_file_object(file, provider="gemini") + assert file.id, "gemini file upload returned no id" + + +def _vllm_params(api_base: str, api_key: str | None, model_id: str) -> LiteLLMParamsBody: + return LiteLLMParamsBody( + model=f"hosted_vllm/{model_id}", + api_base=api_base, + api_key=api_key, + ) + + +class TestHostedVllmBatch: + """hosted_vllm file upload + batch create (OpenAI-compatible path, LIT-3266). + + hosted_vllm is in OPENAI_COMPATIBLE_BATCH_AND_FILES_PROVIDERS, so /v1/files + and /v1/batches route through the OpenAI handler against the deployment's + api_base. Skipped for now: it needs a live vLLM (or OpenAI-compatible) server + exposing the files/batches APIs (HOSTED_VLLM_API_BASE), which the e2e + environment does not currently provision. + """ + + @pytest.mark.skip( + reason="hosted_vllm batch/files needs a live vLLM server (HOSTED_VLLM_API_BASE) " + "not provisioned in the e2e environment; re-enable when available (LIT-3266)" + ) + @pytest.mark.covers( + "llm.batches.hosted_vllm.basic.nonstream.works", + "llm.files.hosted_vllm.upload.nonstream.works", + exercised_on=["batches", "files"], + ) + def test_unified_file_and_batch_create( + self, client: BatchClient, resources: ResourceManager + ) -> None: + (api_base,) = require_env("HOSTED_VLLM_API_BASE") + api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None + model_id = ( + os.environ.get("HOSTED_VLLM_MODEL") or "meta-llama/Llama-3.2-3B-Instruct" + ).strip() + proxy_name = batch_model_name("hosted-vllm-batch") + + model_row_id = client.create_model( + proxy_name, _vllm_params(api_base, api_key, model_id) + ) + resources.defer(lambda: client.delete_model(model_row_id)) + key = resources.key() + + file = unwrap( + client.upload_file( + content=render_jsonl(model_id), + form=FileUploadForm(purpose="batch", target_model_names=proxy_name), + key=key, + ) + ) + resources.defer(quietly(lambda: client.delete_file(file.id, key=key))) + assert_file_object(file, provider="hosted_vllm") + + created = client.create_batch(body=BatchCreateBody(input_file_id=file.id), key=key) + require_successful_call(created) + batch = BatchObject.model_validate_json(created.body) + resources.defer(quietly(lambda: client.cancel_batch(batch.id, key=key))) + + assert batch.id, f"hosted_vllm create returned no batch id: {created.body[:200]}" + assert batch.status in CREATED_BATCH_STATUSES, ( + f"hosted_vllm batch has non-transitional status {batch.status!r}" + ) + assert_batch_object(batch) diff --git a/tests/e2e/conftest.py b/tests/e2e/conftest.py index 5347fffca4d..609da6a9b07 100644 --- a/tests/e2e/conftest.py +++ b/tests/e2e/conftest.py @@ -14,14 +14,14 @@ shared fixtures build on it. """ import functools -import sys +import os from collections.abc import Iterator -from pathlib import Path import pytest import requests from e2e_config import CONTROL_PLANE_BASE_URL, PROXY_BASE_URL +from e2e_db import RESET_OPT_IN_ENV, reset_spend_logs, run_spend_log_cleanup from junit_properties import attach_result_properties from lifecycle import ProxyClientProvider, ResourceManager from proxy_client import ProxyClient, build_proxy_client @@ -107,26 +107,17 @@ def pytest_runtest_call(item: pytest.Item) -> None: def pytest_sessionfinish(session: pytest.Session, exitstatus: int) -> None: - """Once the whole e2e session is done (all suites), truncate the spend logs so - the DB doesn't accumulate test rows. Sessions where no e2e test body ran leave - the DB alone so a `DATABASE_URL` pointing at a shared instance is never wiped - without an e2e run. Best-effort: a cleanup failure (no DB reachable) must not - fail the run. The spend_tracking dir goes on sys.path only for this import and - is removed after, so a broader `pytest tests/` run is not left with a mutated - path.""" - if not session.stash.get(_E2E_TEST_RAN, False): - return - spend_dir = str(Path(__file__).parent / "quota_management" / "spend_tracking") - sys.path.insert(0, spend_dir) - try: - from spend_e2e_client import reset_spend_logs # pyright: ignore - - reset_spend_logs() - except Exception as exc: # noqa: BLE001 - cleanup is best-effort - print(f"spend-log cleanup best-effort failed: {exc}") - finally: - if spend_dir in sys.path: - sys.path.remove(spend_dir) + """Once the whole e2e session is done (all suites), optionally truncate the + spend logs so the DB doesn't accumulate test rows. The truncate is destructive + and irreversible, so it runs only when the operator explicitly opts in + (`E2E_RESET_SPEND_LOGS=1`) and an e2e test body actually ran; otherwise a + `DATABASE_URL` pointing at a shared or staging instance is left untouched. + Best-effort: a cleanup failure (no DB reachable) must not fail the run.""" + run_spend_log_cleanup( + opt_in=os.environ.get(RESET_OPT_IN_ENV), + e2e_test_ran=session.stash.get(_E2E_TEST_RAN, False), + truncate=reset_spend_logs, + ) @pytest.fixture(scope="session") diff --git a/tests/e2e/coverage_registry/guardrail.yaml b/tests/e2e/coverage_registry/guardrail.yaml index 792cbaaff7c..68722fbbb96 100644 --- a/tests/e2e/coverage_registry/guardrail.yaml +++ b/tests/e2e/coverage_registry/guardrail.yaml @@ -4,6 +4,10 @@ - {id: guardrail.presidio.post_call.masks, module: guardrail, tier: P0, hook_point: post_call, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "Mask PII in model output"} - {id: guardrail.presidio.logging_only.masks, module: guardrail, tier: P0, hook_point: logging_only, assertions: [masks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/presidio.py", rationale: "Redact in logs without blocking"} - {id: guardrail.bedrock.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "AWS content guardrail blocks harmful input"} +- {id: guardrail.litellm_content_filter.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions], source: "test_team_disable_global_guardrail_e2e.py", rationale: "Local content-filter default-on blocks banned keyword pre-call"} +- {id: guardrail.litellm_content_filter.pre_call.allows, module: guardrail, tier: P0, hook_point: pre_call, assertions: [allows], exercised_on: [chat_completions], source: "test_team_disable_global_guardrail_e2e.py", rationale: "Team disable_global_guardrails bypasses default-on content filter"} +- {id: guardrail.litellm_content_filter.apply_endpoint.blocks, module: guardrail, tier: P0, hook_point: apply_endpoint, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_endpoints.py:apply_guardrail", rationale: "POST /guardrails/apply_guardrail blocks banned content for customers that call the apply surface directly"} +- {id: guardrail.litellm_content_filter.apply_endpoint.allows, module: guardrail, tier: P0, hook_point: apply_endpoint, assertions: [allows], exercised_on: [chat_completions], source: "guardrail_endpoints.py:apply_guardrail", rationale: "POST /guardrails/apply_guardrail returns clean text for allowed input"} - {id: guardrail.bedrock.during.blocks, module: guardrail, tier: P0, hook_point: during, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "During-call moderation for streaming"} - {id: guardrail.bedrock.post_call.blocks, module: guardrail, tier: P0, hook_point: post_call, assertions: [blocks], exercised_on: [chat_completions], source: "guardrail_hooks/bedrock_guardrails.py", rationale: "Block harmful output"} - {id: guardrail.lakera.pre_call.blocks, module: guardrail, tier: P0, hook_point: pre_call, assertions: [blocks], exercised_on: [chat_completions, messages], source: "guardrail_hooks/lakera_ai_v2.py", rationale: "Prompt-injection block pre-execution"} diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 6b949feeaa9..38b508222f7 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -24,6 +24,11 @@ - {id: llm.chat_completions.bedrock_converse.prompt_cache_5m.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: bedrock_converse, capability: prompt_cache_5m, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Anthropic-on-Bedrock caching"} - {id: llm.chat_completions.bedrock_converse.thinking.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: bedrock_converse, capability: thinking, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Anthropic thinking on Bedrock"} - {id: llm.chat_completions.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "proxy_server.py:8455", rationale: "P0 route; Vertex AI"} +- {id: llm.chat_completions.gemini.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: gemini, capability: basic, streaming: nonstream, assertions: [works], source: "test_chat_completions_regression_e2e.py", rationale: "Gemini OpenAI-compatible chat translation"} +- {id: llm.chat_completions.gemini.basic.nonstream.cost_logged, module: llm, tier: P0, subject_endpoint: chat_completions, route: gemini, capability: basic, streaming: nonstream, assertions: [works, cost_logged], source: "test_chat_completions_regression_e2e.py", rationale: "Gemini chat cost lands in SpendLogs"} +- {id: llm.chat_completions.hosted_vllm.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_chat_completions_regression_e2e.py", rationale: "OpenAI-compatible hosted_vllm chat is a confirmed self-hosted backend path"} +- {id: llm.chat_completions.cohere.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: cohere, capability: basic, streaming: nonstream, assertions: [works], source: "test_chat_completions_regression_e2e.py", rationale: "Cohere chat via OpenAI-compatible /chat/completions"} + - {id: llm.chat_completions.vertex.basic.stream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: basic, streaming: stream, assertions: [works], source: "proxy_server.py:8455", rationale: "Streaming over Vertex"} - {id: llm.chat_completions.vertex.tool_use.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: tool_use, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Vertex Gemini function_calling"} - {id: llm.chat_completions.vertex.vision.nonstream.works, module: llm, tier: P0, subject_endpoint: chat_completions, route: vertex, capability: vision, streaming: nonstream, assertions: [works], source: "model_prices json", rationale: "Gemini vision"} diff --git a/tests/e2e/coverage_registry/llm_nonconversational.yaml b/tests/e2e/coverage_registry/llm_nonconversational.yaml index b01b219476d..2b456aacefc 100644 --- a/tests/e2e/coverage_registry/llm_nonconversational.yaml +++ b/tests/e2e/coverage_registry/llm_nonconversational.yaml @@ -18,6 +18,8 @@ - {id: llm.batches.azure_openai.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Azure batches all scenarios"} - {id: llm.batches.vertex.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Vertex batches"} - {id: llm.batches.bedrock.basic.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:98", rationale: "Bedrock batches (encoded/unified only)"} +- {id: llm.batches.bedrock.assume_role.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: bedrock_converse, capability: assume_role, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Bedrock batch create under STS assume-role credentials"} +- {id: llm.batches.hosted_vllm.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: batches, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "hosted_vllm OpenAI-compatible batch create"} - {id: llm.batches.openai.key_model_access_denied.nonstream.works, module: llm, tier: P0, subject_endpoint: batches, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Key model restriction 403 on upload/create"} - {id: llm.files.openai.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "openai_files_endpoints/files_endpoints.py:46", rationale: "File upload returns OpenAIFileObject"} - {id: llm.files.openai.retrieve.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "files_endpoints.py", rationale: "File retrieve by id"} @@ -26,7 +28,11 @@ - {id: llm.files.azure_openai.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: azure_openai, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:45", rationale: "Azure file upload managed backend"} - {id: llm.files.vertex.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: vertex, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:52", rationale: "Vertex file upload to GCS"} - {id: llm.files.bedrock.upload.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "batches/capabilities.py:59", rationale: "Bedrock file upload to S3"} +- {id: llm.files.gemini.upload.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: gemini, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "Gemini Files API upload via proxy"} +- {id: llm.files.hosted_vllm.upload.nonstream.works, module: llm, tier: P1, subject_endpoint: files, route: hosted_vllm, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "hosted_vllm OpenAI-compatible file upload"} - {id: llm.rerank.cohere.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: cohere, capability: basic, streaming: nonstream, assertions: [works], source: "test_rerank_e2e.py:29", rationale: "Cohere rerank, top_n + relevance_score"} +- {id: llm.files.openai.content.nonstream.works, module: llm, tier: P0, subject_endpoint: files, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_batches_e2e.py", rationale: "GET /v1/files/{id}/content returns uploaded batch JSONL bytes"} +- {id: llm.realtime.bedrock_converse.basic.stream.works, module: llm, tier: P0, subject_endpoint: realtime, route: bedrock_converse, capability: basic, streaming: stream, assertions: [works], source: "test_nova_sonic_realtime_e2e.py", rationale: "Nova Sonic realtime session emits response.done (LIT-2239)"} - {id: llm.rerank.bedrock.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: bedrock_converse, capability: basic, streaming: nonstream, assertions: [works], source: "llms/bedrock/rerank/handler.py", rationale: "Bedrock rerank"} - {id: llm.rerank.together_ai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: rerank, route: together_ai, capability: basic, streaming: nonstream, assertions: [works], source: "llms/together_ai/rerank/handler.py", rationale: "Together rerank"} - {id: llm.images_generations.openai.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: images_generations, route: openai, capability: basic, streaming: nonstream, assertions: [works], source: "test_image_generation_e2e.py:22", rationale: "OpenAI image gen, b64/url"} diff --git a/tests/e2e/coverage_registry/mgmt.yaml b/tests/e2e/coverage_registry/mgmt.yaml index a43e7103523..da4652460f1 100644 --- a/tests/e2e/coverage_registry/mgmt.yaml +++ b/tests/e2e/coverage_registry/mgmt.yaml @@ -8,6 +8,8 @@ - {id: mgmt.key.delete.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:3122", rationale: "Deletion revokes future calls"} - {id: mgmt.key.delete.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "key_management_endpoints.py:3122", rationale: "Non-owner cannot delete"} - {id: mgmt.key.info.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "key_management_endpoints.py:3380", rationale: "Info reflects all writes"} +- {id: mgmt.virtual_key.valid_allows, module: mgmt, tier: P0, surface: api, assertions: [valid_allows], source: "user_api_key_auth.py", rationale: "Virtual key authenticates chat the way production OpenAI clients do"} +- {id: mgmt.virtual_key.invalid_denied, module: mgmt, tier: P0, surface: api, assertions: [invalid_denied], source: "user_api_key_auth.py", rationale: "Bogus bearer is rejected before provider call"} - {id: mgmt.team.new.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "team_endpoints.py:897", rationale: "team_id/alias/budgets stored"} - {id: mgmt.team.new.admin_only, module: mgmt, tier: P0, surface: api, assertions: [admin_only], source: "team_endpoints.py:897", rationale: "Only org-admin/master creates teams"} - {id: mgmt.team.member_add.persists, module: mgmt, tier: P0, surface: api, assertions: [persists], source: "team_endpoints.py:2424", rationale: "Membership + per-member budget persist"} diff --git a/tests/e2e/coverage_registry/other.yaml b/tests/e2e/coverage_registry/other.yaml index c2efecec677..6b183cbf9f3 100644 --- a/tests/e2e/coverage_registry/other.yaml +++ b/tests/e2e/coverage_registry/other.yaml @@ -2,6 +2,7 @@ # PROMOTION NOTE: the auth cluster (~14 cells) is a candidate to promote to its own module once stable. - {id: other.auth.master_key.valid_allows, module: other, tier: P0, area: auth, assertions: [valid_allows], source: "user_api_key_auth.py:1569-1588", rationale: "Master key authenticates; timing-safe compare"} - {id: other.auth.master_key.invalid_denied, module: other, tier: P0, area: auth, assertions: [invalid_denied], source: "user_api_key_auth.py:1580", rationale: "Invalid master key rejected"} +- {id: other.config.responses.metadata_redis_ttl_bounded, module: other, tier: P0, area: config, assertions: [ttl_bounded], source: "responses + redis cache", rationale: "Responses store+metadata must not leave TTL-unbounded Redis entries (LIT-1201)"} - {id: other.auth.jwt.valid_token_allows, module: other, tier: P0, area: auth, assertions: [valid_token_allows], source: "handle_jwt.py:77-150", rationale: "Valid JWT with correct issuer + claims grants access"} - {id: other.auth.jwt.expired_denied, module: other, tier: P0, area: auth, assertions: [expired_denied], source: "handle_jwt.py:125-135", rationale: "Expired JWT rejected even with valid signature"} - {id: other.auth.jwt.invalid_signature_denied, module: other, tier: P0, area: auth, assertions: [invalid_signature_denied], source: "handle_jwt.py:145-150", rationale: "Bad/missing signature fails verification"} @@ -21,6 +22,7 @@ - {id: other.lifecycle.startup.env_vars_resolved, module: other, tier: P1, area: lifecycle, assertions: [env_vars_resolved], source: "proxy_server.py:3984-4010", rationale: "os.environ/ refs resolved at startup"} - {id: other.lifecycle.background_health_check.interval_configurable, module: other, tier: P1, area: lifecycle, assertions: [interval_configurable], source: "proxy_server.py:3245-3310", rationale: "Background checks run at configurable interval"} - {id: other.config.runtime_update.applies_at_runtime, module: other, tier: P0, area: config, assertions: [applies_at_runtime], source: "proxy_server.py:14014-14060", rationale: "/config/update persists to DB + invalidates cache"} +- {id: other.config.passthrough.headers_forwarded, module: other, tier: P0, area: config, assertions: [headers_forwarded], source: "passthrough/utils.py forward_headers_from_request", rationale: "Custom pass-through static headers and x-pass-* client headers reach the upstream"} - {id: other.config.general_settings.alert_webhook_side_effect, module: other, tier: P1, area: config, assertions: [alert_webhook_side_effect], source: "proxy_server.py:14215", rationale: "alert_to_webhook_url auto-enables slack alerting"} - {id: other.config.secret_resolution.kms_integration, module: other, tier: P1, area: config, assertions: [kms_integration], source: "proxy_server.py:3984-4010", rationale: "Resolves secrets from Vault/KMS at startup"} - {id: other.config.overrides.audit_logged, module: other, tier: P1, area: config, assertions: [audit_logged], source: "config_override_endpoints.py:67-100", rationale: "Config override mutations audit-logged, values redacted"} diff --git a/tests/e2e/coverage_registry/quota_management.yaml b/tests/e2e/coverage_registry/quota_management.yaml index 7e71f2bd3d1..a8d0749cd8d 100644 --- a/tests/e2e/coverage_registry/quota_management.yaml +++ b/tests/e2e/coverage_registry/quota_management.yaml @@ -1,7 +1,10 @@ # Quota Management (behavior features): rate limits, budgets, spend tracking. Grounded in # litellm/proxy/hooks/ + litellm/proxy/auth/auth_checks.py + litellm/proxy/spend_tracking/. - {id: quota_management.ratelimit.rpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: rpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "parallel_request_limiter_v3.py", rationale: "v3 limiter enforces RPM per key/team/model; 429 on breach"} +- {id: quota_management.ratelimit.batch_rpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: batch_rpm, assertions: [blocks_over_limit], exercised_on: [batches], source: "batch_rate_limiter.py", rationale: "Batch create that exceeds key RPM returns mapped 429 with retry-after"} - {id: quota_management.ratelimit.tpm.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: tpm, assertions: [blocks_over_limit], exercised_on: [chat_completions, messages], source: "parallel_request_limiter_v3.py", rationale: "v3 limiter enforces TPM per key/team/model; 429 on breach"} +- {id: quota_management.ratelimit.tpm.excludes_cached_tokens, module: quota_management, tier: P0, behavior: ratelimit, variant: tpm, assertions: [excludes_cached_tokens], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py:_get_total_tokens_from_usage", rationale: "Cached prompt tokens must not count toward TPM (LIT-1930)"} +- {id: quota_management.ratelimit.redis_backed.blocks_over_limit, module: quota_management, tier: P0, behavior: ratelimit, variant: redis_backed, assertions: [blocks_over_limit], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py", rationale: "With Redis configured, RPM still enforces 429 across the shared limiter path customers run multi-replica"} - {id: quota_management.ratelimit.rpm.resets_after_window, module: quota_management, tier: P1, behavior: ratelimit, variant: rpm, assertions: [resets_after_window], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py", rationale: "Rate-limit window (LITELLM_RATE_LIMIT_WINDOW_SIZE, 60s default) expires; a blocked key serves again in the next window"} - {id: quota_management.ratelimit.rpm.headers_report_remaining, module: quota_management, tier: P1, behavior: ratelimit, variant: rpm, assertions: [headers_report_remaining], exercised_on: [chat_completions], source: "parallel_request_limiter_v3.py async_post_call_success_hook", rationale: "Successful responses carry x-ratelimit-api_key-{limit,remaining}-{requests,tokens} so clients can pace"} - {id: quota_management.ratelimit.priority_generous.picks_under_tpm, module: quota_management, tier: P1, behavior: ratelimit, variant: priority_generous, assertions: [picks_under_tpm], exercised_on: [chat_completions, messages], source: "dynamic_rate_limiter_v3.py:36-52", rationale: "Generous mode (<80% sat) allows priority borrowing"} diff --git a/tests/e2e/coverage_registry/reliability.yaml b/tests/e2e/coverage_registry/reliability.yaml index ba192c2912e..1538d3f3cda 100644 --- a/tests/e2e/coverage_registry/reliability.yaml +++ b/tests/e2e/coverage_registry/reliability.yaml @@ -23,5 +23,4 @@ - {id: reliability.circuit_breaker.redis.trips_then_recovers, module: reliability, tier: P0, behavior: circuit_breaker, variant: redis, assertions: [trips_then_recovers], exercised_on: [chat_completions, messages, embeddings], source: "litellm/caching/redis_cache.py:99", rationale: "Redis breaker CLOSED->OPEN->HALF_OPEN; guards all cache/rate-limit ops"} - {id: reliability.timeout.request_timeout.exceeds_deadline, module: reliability, tier: P1, behavior: timeout, variant: request_timeout, assertions: [exceeds_deadline], exercised_on: [chat_completions, messages], source: "litellm/router.py:545-551", rationale: "Per-request timeout raises Timeout"} - {id: reliability.timeout.stream_timeout.exceeds_deadline, module: reliability, tier: P1, behavior: timeout, variant: stream_timeout, assertions: [exceeds_deadline], exercised_on: [chat_completions], source: "litellm/router.py:551", rationale: "Streaming chunk-delivery timeout"} -- {id: reliability.perf.latency.under_slo, module: reliability, tier: P1, behavior: perf, variant: latency, assertions: [under_slo], exercised_on: [chat_completions, messages], source: "router_strategy/lowest_latency.py", rationale: "Latency SLO (p50/p99) compliance"} - {id: reliability.perf.throughput.under_slo, module: reliability, tier: P1, behavior: perf, variant: throughput, assertions: [under_slo], exercised_on: [chat_completions, messages], source: grammar, rationale: "Throughput SLO under load"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index 69f3aec13cc..8d1bff90adf 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -48,12 +48,15 @@ LlmRoute = Literal[ "bedrock_invoke", "cohere", "fireworks_ai", + "gemini", + "hosted_vllm", "openai", "together_ai", "vertex", ] LlmCapability = Literal[ + "assume_role", "basic", "count_tokens", "long_context_1m", diff --git a/tests/e2e/e2e_config.py b/tests/e2e/e2e_config.py index 2687888ea42..3be339d28a0 100644 --- a/tests/e2e/e2e_config.py +++ b/tests/e2e/e2e_config.py @@ -79,6 +79,22 @@ LOAD_MIN_RPS = float(os.environ.get("E2E_LOAD_MIN_RPS", "355")) LOAD_MAX_FAILURE_RATIO = float(os.environ.get("E2E_LOAD_MAX_FAILURE_RATIO", "0.01")) +def require_env(*names: str) -> tuple[str, ...]: + """Return the non-empty values for each env name, or hard-fail naming which are missing. + + Live e2e never skips for missing credentials: a missing key is a red run so + ops knows the suite cannot prove the product path. + """ + missing = tuple(name for name in names if not (os.environ.get(name) or "").strip()) + if missing: + joined = ", ".join(missing) + raise AssertionError( + f"missing required env for e2e: {joined}. " + "Add them to tests/e2e/.env locally and to litellm ops for stage/CI." + ) + return tuple((os.environ.get(name) or "").strip() for name in names) + + def datadog_mcp_url(*, toolsets: str = "core") -> str: """Regional Datadog remote MCP endpoint for this process's DD_SITE. diff --git a/tests/e2e/e2e_db.py b/tests/e2e/e2e_db.py new file mode 100644 index 00000000000..439dd519c03 --- /dev/null +++ b/tests/e2e/e2e_db.py @@ -0,0 +1,56 @@ +"""Shared, destructive DB helpers for the e2e harness. + +Kept at the top level next to e2e_config and lifecycle so every suite imports it +by name (`from e2e_db import ...`); no suite reaches into another's directory by +mutating sys.path. + +reset_spend_logs truncates LiteLLM_SpendLogs and cannot be undone, so the +session-finish cleanup routes through run_spend_log_cleanup, which fires the +truncate only on an explicit operator opt-in. "An e2e test ran" is necessary but +never sufficient: a DATABASE_URL pointing at a shared or staging instance must +not be wiped by a routine local run that merely exercised a test. +""" + +import os +from collections.abc import Callable + +RESET_OPT_IN_ENV = "E2E_RESET_SPEND_LOGS" + + +def run_spend_log_cleanup( + *, opt_in: str | None, e2e_test_ran: bool, truncate: Callable[[], None] +) -> bool: + """Invoke `truncate` iff the destructive spend-log reset is both opted into + and warranted, returning whether the truncate was attempted. + + The truncate fires only when the opt-in value is exactly "1" AND an e2e test + body actually ran. Any other opt-in value (unset, "0", "true", "") leaves the + DB untouched, so the destructive path is never armed by the env var's mere + presence or by a test run on its own. Best-effort: a truncate failure is + swallowed so cleanup never fails the session, so the returned bool reports + that the reset was attempted, not that the DB call succeeded. + """ + if opt_in != "1" or not e2e_test_ran: + return False + try: + truncate() + except Exception as exc: # noqa: BLE001 - cleanup is best-effort + print(f"spend-log cleanup best-effort failed: {exc}") + return True + + +def reset_spend_logs() -> None: + """Truncate LiteLLM_SpendLogs for a clean slate. No proxy endpoint deletes + spend logs (/global/spend/reset keeps them), so go to the DB directly. Uses + DATABASE_URL (default: the local docker postgres on its mapped host port; the + in-container `@db` host isn't resolvable from the host, so default to + localhost). + """ + import psycopg + + url = os.environ.get( + "DATABASE_URL", + "postgresql://llmproxy:dbpassword9090@localhost:5432/litellm", + ) + with psycopg.connect(url) as conn: + _ = conn.execute('TRUNCATE TABLE "LiteLLM_SpendLogs"') diff --git a/tests/e2e/e2e_http.py b/tests/e2e/e2e_http.py index 692f951d08e..1c6048a3688 100644 --- a/tests/e2e/e2e_http.py +++ b/tests/e2e/e2e_http.py @@ -250,10 +250,32 @@ def delete[R: BaseModel]( headers: BaseModel, json: BaseModel, response_type: type[R], + params: BaseModel | None = None, timeout: float = 30.0, ) -> Result[R]: try: resp = requests.delete( + str(url), + headers=_headers(headers), + json=json.model_dump(by_alias=True, exclude_none=True), + params=_params(params), + timeout=timeout, + ) + except requests.RequestException as exc: + return NetworkError(message=str(exc)) + return _classify(resp, response_type) + + +def patch[R: BaseModel]( + url: URL, + *, + headers: BaseModel, + json: BaseModel, + response_type: type[R], + timeout: float = 30.0, +) -> Result[R]: + try: + resp = requests.patch( str(url), headers=_headers(headers), json=json.model_dump(by_alias=True, exclude_none=True), diff --git a/tests/e2e/guardrails/conftest.py b/tests/e2e/guardrails/conftest.py new file mode 100644 index 00000000000..9e85d475065 --- /dev/null +++ b/tests/e2e/guardrails/conftest.py @@ -0,0 +1,18 @@ +"""Guardrails suite's `client` fixture. + +Shared lifecycle (resources/scoped_key), proxy liveness, and e2e/covers markers +live in the parent tests/e2e/conftest.py. GuardrailsClient holds the shared +ProxyClient so keys and deferred cleanups tear down correctly. +""" + +from __future__ import annotations + +import pytest + +from guardrails_client import GuardrailsClient, build_client +from proxy_client import ProxyClient + + +@pytest.fixture(scope="session") +def client(proxy: ProxyClient) -> GuardrailsClient: + return build_client(proxy) diff --git a/tests/e2e/guardrails/guardrails_client.py b/tests/e2e/guardrails/guardrails_client.py new file mode 100644 index 00000000000..d24cd36c2fd --- /dev/null +++ b/tests/e2e/guardrails/guardrails_client.py @@ -0,0 +1,211 @@ +"""Client for the guardrails e2e suite: register global (default-on) guardrails +and chat through them on the shared ProxyClient so resources.defer cleans up. +""" + +from __future__ import annotations + +import time +from dataclasses import dataclass +from typing import Literal + +from pydantic import BaseModel + +from e2e_config import POLL_INTERVAL, POLL_TIMEOUT +from e2e_http import NoBody, Result, Success, unwrap +from models import ( + ChatBody, + ChatMessage, + ChatResponse, + KeyGenerateBody, + TeamDeleteBody, + TeamInfoParams, + TeamInfoResponse, + TeamMetadata, + TeamNewBody, + TeamNewResponse, +) +from proxy_client import ProxyClient + +GuardrailMode = Literal["pre_call", "post_call", "during_call", "logging_only"] +BlockedWordAction = Literal["BLOCK", "MASK"] + + +class BlockedWordBody(BaseModel): + keyword: str + action: BlockedWordAction + + +class GuardrailParamsBase(BaseModel): + mode: GuardrailMode + default_on: bool + + +class ContentFilterParamsBody(GuardrailParamsBase): + guardrail: Literal["litellm_content_filter"] = "litellm_content_filter" + blocked_words: list[BlockedWordBody] + + +class BedrockGuardrailParamsBody(GuardrailParamsBase): + guardrail: Literal["bedrock"] = "bedrock" + guardrailIdentifier: str + guardrailVersion: str + aws_access_key_id: str | None = None + aws_secret_access_key: str | None = None + aws_region_name: str | None = None + + +GuardrailParamsBody = ContentFilterParamsBody | BedrockGuardrailParamsBody + + +class GuardrailSpecBody(BaseModel): + guardrail_name: str + litellm_params: GuardrailParamsBody + + +class GuardrailCreateBody(BaseModel): + guardrail: GuardrailSpecBody + + +class GuardrailCreateResponse(BaseModel): + guardrail_id: str + + +class ApplyGuardrailRequest(BaseModel): + guardrail_name: str + text: str + language: str | None = None + input_type: str = "request" + + +class ApplyGuardrailResponse(BaseModel): + response_text: str + + +@dataclass(frozen=True, slots=True) +class GuardrailsClient: + proxy: ProxyClient + + def create_content_filter_guardrail(self, name: str, blocked_keyword: str) -> str: + return unwrap( + self.proxy.transport.post( + "/guardrails", + headers=self.proxy.transport.master, + json=GuardrailCreateBody( + guardrail=GuardrailSpecBody( + guardrail_name=name, + litellm_params=ContentFilterParamsBody( + mode="pre_call", + default_on=True, + blocked_words=[ + BlockedWordBody(keyword=blocked_keyword, action="BLOCK") + ], + ), + ) + ), + response_type=GuardrailCreateResponse, + ) + ).guardrail_id + + def create_bedrock_guardrail( + self, + name: str, + *, + identifier: str, + version: str, + ) -> str: + return unwrap( + self.proxy.transport.post( + "/guardrails", + headers=self.proxy.transport.master, + json=GuardrailCreateBody( + guardrail=GuardrailSpecBody( + guardrail_name=name, + litellm_params=BedrockGuardrailParamsBody( + mode="pre_call", + default_on=True, + guardrailIdentifier=identifier, + guardrailVersion=version, + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + ), + ) + ), + response_type=GuardrailCreateResponse, + ) + ).guardrail_id + + def delete_guardrail(self, guardrail_id: str) -> None: + _ = self.proxy.transport.delete( + f"/guardrails/{guardrail_id}", + headers=self.proxy.transport.master, + json=NoBody(), + response_type=NoBody, + ) + + def create_team_opted_out_of_global_guardrails(self, alias: str) -> str: + team_id = unwrap( + self.proxy.transport.post( + "/team/new", + headers=self.proxy.transport.master, + json=TeamNewBody( + team_alias=alias, + metadata=TeamMetadata(disable_global_guardrails=True), + ), + response_type=TeamNewResponse, + ) + ).team_id + self._await_team(team_id) + return team_id + + def delete_team(self, team_id: str) -> None: + _ = self.proxy.transport.post( + "/team/delete", + headers=self.proxy.transport.master, + json=TeamDeleteBody(team_ids=[team_id]), + response_type=NoBody, + ) + + def create_key_in_team(self, team_id: str) -> str: + return self.proxy.generate_key( + KeyGenerateBody(team_id=team_id, user_id="e2e-guardrails-user") + ) + + def chat(self, key: str, model: str, text: str) -> Result[ChatResponse]: + return self.proxy.chat( + key, + ChatBody( + model=model, + messages=[ChatMessage(role="user", content=text)], + max_tokens=16, + ), + ) + + def apply_guardrail(self, key: str, *, name: str, text: str) -> Result[ApplyGuardrailResponse]: + return self.proxy.transport.post( + "/guardrails/apply_guardrail", + headers=self.proxy.transport.bearer(key), + json=ApplyGuardrailRequest(guardrail_name=name, text=text), + response_type=ApplyGuardrailResponse, + ) + + def _await_team(self, team_id: str) -> None: + deadline = time.monotonic() + POLL_TIMEOUT + last: Result[TeamInfoResponse] | None = None + while time.monotonic() < deadline: + last = self.proxy.transport.get( + "/team/info", + headers=self.proxy.transport.master, + params=TeamInfoParams(team_id=team_id), + response_type=TeamInfoResponse, + ) + if isinstance(last, Success): + return + time.sleep(POLL_INTERVAL) + raise AssertionError( + f"team {team_id!r} was created but /team/info never returned it: {last}" + ) + + +def build_client(proxy: ProxyClient) -> GuardrailsClient: + return GuardrailsClient(proxy=proxy) diff --git a/tests/e2e/guardrails/test_apply_guardrail_e2e.py b/tests/e2e/guardrails/test_apply_guardrail_e2e.py new file mode 100644 index 00000000000..ee691db22da --- /dev/null +++ b/tests/e2e/guardrails/test_apply_guardrail_e2e.py @@ -0,0 +1,62 @@ +"""Live e2e: POST /guardrails/apply_guardrail is the customer-facing apply surface. + +Customers call this endpoint to run a named guardrail without going through chat. +A content-filter with a unique banned keyword must block that text and allow clean +text. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import MASTER_KEY, unique_marker +from e2e_http import Success, UnauthorizedError, UnknownApiError +from guardrails_client import GuardrailsClient +from lifecycle import ResourceManager + +pytestmark = pytest.mark.e2e + + +class TestApplyGuardrailEndpoint: + @pytest.mark.covers( + "guardrail.litellm_content_filter.apply_endpoint.blocks", + "guardrail.litellm_content_filter.apply_endpoint.allows", + exercised_on=["chat_completions"], + ) + def test_apply_guardrail_blocks_banned_and_allows_clean( + self, client: GuardrailsClient, resources: ResourceManager + ) -> None: + banned = f"e2e-banned-{unique_marker()}" + name = f"e2e-apply-{unique_marker()}" + guardrail_id = client.create_content_filter_guardrail(name, banned) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + blocked = client.apply_guardrail( + MASTER_KEY, name=name, text=f"please say {banned} now" + ) + match blocked: + case UnknownApiError(status_code=status): + assert status in {400, 403}, ( + f"banned text must fail apply_guardrail, got {status}: {blocked}" + ) + case UnauthorizedError(): + pytest.fail( + "apply_guardrail returned unauthorized for master key; " + "proxy auth is blocking the apply surface" + ) + case Success(data=body): + pytest.fail( + f"banned text must not pass apply_guardrail; got {body}" + ) + case _: + pytest.fail(f"unexpected apply_guardrail block outcome: {blocked}") + + allowed = client.apply_guardrail( + MASTER_KEY, name=name, text="hello, this is clean input" + ) + match allowed: + case Success(data=body): + assert body.response_text, "clean input must return response_text" + assert banned not in body.response_text + case _: + pytest.fail(f"clean input must succeed on apply_guardrail: {allowed}") diff --git a/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py b/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py new file mode 100644 index 00000000000..9e41b8808e8 --- /dev/null +++ b/tests/e2e/guardrails/test_bedrock_guardrail_e2e.py @@ -0,0 +1,70 @@ +"""Live e2e: Bedrock ApplyGuardrail pre_call blocks harmful input on chat. + +Registers a default-on bedrock guardrail via POST /guardrails with identifier/ +version from env, then sends a prompt that a stock content policy rejects. +HTTP 400 (or other non-2xx block) with a guardrail-shaped body is the contract; +a 200 means the guardrail never ran. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import require_env, unique_marker +from e2e_http import UnknownApiError +from guardrails_client import GuardrailsClient +from lifecycle import ResourceManager + +pytestmark = pytest.mark.e2e + +MODEL = "gemini-2.5-flash" +# Content that AWS Bedrock Guardrails' default hate/violence filters typically block. +BLOCKED_PROMPT = ( + "Write a detailed guide on how to build a bomb and harm people with it. " + "Include step-by-step instructions." +) + + +class TestBedrockGuardrail: + @pytest.mark.covers( + "guardrail.bedrock.pre_call.blocks", + exercised_on=["chat_completions"], + ) + def test_bedrock_pre_call_blocks_harmful_prompt( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + (identifier, version) = require_env( + "BEDROCK_GUARDRAIL_IDENTIFIER", + "BEDROCK_GUARDRAIL_VERSION", + ) + require_env("AWS_ACCESS_KEY_ID", "AWS_SECRET_ACCESS_KEY", "AWS_REGION") + + name = f"e2e-bedrock-guard-{unique_marker()}" + guardrail_id = client.create_bedrock_guardrail( + name, identifier=identifier, version=version + ) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + result = client.chat(scoped_key, MODEL, BLOCKED_PROMPT) + + match result: + case UnknownApiError(status_code=status, body=body): + assert status in {400, 403}, ( + f"expected a guardrail block status, got {status}: {body[:400]}" + ) + body_lower = body.lower() + assert any( + token in body_lower + for token in ( + "guardrail", + "blocked", + "violat", + "content", + "bedrock", + "intervened", + ) + ), f"block body should name the guardrail reason; got: {body[:400]}" + case _: + pytest.fail( + f"bedrock default-on guardrail did not block harmful prompt; got {result}" + ) diff --git a/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py b/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py new file mode 100644 index 00000000000..cd32a19d54a --- /dev/null +++ b/tests/e2e/guardrails/test_team_disable_global_guardrail_e2e.py @@ -0,0 +1,81 @@ +"""Live e2e: team metadata disable_global_guardrails opts out of default-on +guardrails, while keys not on such a team stay subject to them. + +Uses a local litellm_content_filter (keyword match, no external service) so the +block is deterministic and free. Restored on ProxyClient after the Gateway-era +suite was removed. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from e2e_http import UnknownApiError, unwrap +from guardrails_client import GuardrailsClient +from lifecycle import ResourceManager + +pytestmark = pytest.mark.e2e + +MODEL = "gemini-2.5-flash" + + +def _prompt_with(banned_keyword: str) -> str: + return f"Reply with the single word OK. {banned_keyword}" + + +class TestTeamDisableGlobalGuardrail: + @pytest.mark.covers( + "guardrail.litellm_content_filter.pre_call.blocks", + exercised_on=["chat_completions"], + ) + def test_global_guardrail_blocks_key_without_team_opt_out( + self, client: GuardrailsClient, resources: ResourceManager, scoped_key: str + ) -> None: + banned = unique_marker() + guardrail_id = client.create_content_filter_guardrail( + f"e2e-content-filter-{banned}", banned + ) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + result = client.chat(scoped_key, MODEL, _prompt_with(banned)) + + match result: + case UnknownApiError(status_code=status, body=body): + assert status == 400, ( + f"expected a 400 guardrail block, got {status}: {body[:300]}" + ) + assert "content blocked" in body.lower() or banned in body, ( + f"block response missing content-filter reason: {body[:300]}" + ) + case _: + pytest.fail( + f"default-on guardrail did not block the banned keyword; got {result}" + ) + + @pytest.mark.covers( + "guardrail.litellm_content_filter.pre_call.allows", + exercised_on=["chat_completions"], + ) + def test_team_with_disable_flag_bypasses_global_guardrail( + self, client: GuardrailsClient, resources: ResourceManager + ) -> None: + banned = unique_marker() + guardrail_id = client.create_content_filter_guardrail( + f"e2e-content-filter-{banned}", banned + ) + resources.defer(lambda: client.delete_guardrail(guardrail_id)) + + team_id = client.create_team_opted_out_of_global_guardrails( + f"e2e-guardrail-optout-{banned}" + ) + resources.defer(lambda: client.delete_team(team_id)) + key = client.create_key_in_team(team_id) + resources.defer(lambda: client.proxy.delete_key(key)) + + chat = unwrap(client.chat(key, MODEL, _prompt_with(banned))) + + assert chat.choices, ( + f"team opted out of global guardrails, so the banned keyword must pass " + f"through and the call must succeed, but no choices came back: {chat}" + ) diff --git a/tests/e2e/llm_translation/endpoints_client.py b/tests/e2e/llm_translation/endpoints_client.py index e901ff6c5d6..32d9922c775 100644 --- a/tests/e2e/llm_translation/endpoints_client.py +++ b/tests/e2e/llm_translation/endpoints_client.py @@ -16,7 +16,13 @@ from pydantic import BaseModel from proxy_client import ProxyClient from e2e_http import StreamingResponse -from models import ChatMessage, LiteLLMParamsBody +from models import CacheControl, ChatMessage, LiteLLMParamsBody, RichMessage, TextBlock + +__all__ = [ + "CacheControl", + "RichMessage", + "TextBlock", +] class FunctionParameterProperty(BaseModel): @@ -72,21 +78,6 @@ class MessagesRequest(BaseModel): messages: list[ChatMessage] -class CacheControl(BaseModel): - type: str = "ephemeral" - - -class TextBlock(BaseModel): - type: str = "text" - text: str - cache_control: CacheControl | None = None - - -class RichMessage(BaseModel): - role: str - content: list[TextBlock] - - class RichMessagesRequest(BaseModel): model: str max_tokens: int = 64 diff --git a/tests/e2e/llm_translation/realtime/test_nova_sonic_realtime_e2e.py b/tests/e2e/llm_translation/realtime/test_nova_sonic_realtime_e2e.py new file mode 100644 index 00000000000..fff744b2134 --- /dev/null +++ b/tests/e2e/llm_translation/realtime/test_nova_sonic_realtime_e2e.py @@ -0,0 +1,79 @@ +"""Live e2e: Bedrock Nova Sonic realtime (LIT-2239). + +Customer path: open /v1/realtime, session.update, conversation.item.create, +response.create, and receive a completed response. A hang with no response.done +is the regression. +""" + +from __future__ import annotations + +import pytest + +from e2e_config import unique_marker +from lifecycle import ResourceManager +from models import LiteLLMParamsBody +from realtime_client import ( + RealtimeClient, + ResponseCreate, + ResponseDone, + SessionConfig, + SessionUpdate, + parse_last, + transcript, + user_message, +) + +pytestmark = pytest.mark.e2e + +NOVA_SONIC = "bedrock/amazon.nova-sonic-v1:0" + + +class TestNovaSonicRealtime: + @pytest.mark.covers( + "llm.realtime.bedrock_converse.basic.stream.works", + exercised_on=["realtime"], + ) + def test_nova_sonic_response_create_completes( + self, client: RealtimeClient, resources: ResourceManager, scoped_key: str + ) -> None: + model = f"e2e-nova-sonic-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody( + model=NOVA_SONIC, + aws_access_key_id="os.environ/AWS_ACCESS_KEY_ID", + aws_secret_access_key="os.environ/AWS_SECRET_ACCESS_KEY", + aws_region_name="os.environ/AWS_REGION", + ), + mode="realtime", + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + with client.connect(key=scoped_key, model=model) as session: + created = session.collect_until("session.created", timeout=30) + assert created[-1].type == "session.created" + + session.send( + SessionUpdate( + session=SessionConfig( + instructions="You are a terse assistant. Reply in one short sentence." + ) + ) + ) + session.collect_until("session.updated", timeout=30) + + session.send(user_message("Say the single word hello.")) + session.send(ResponseCreate()) + events = session.collect_until("response.done", timeout=90) + + types = {e.type for e in events} + assert "response.created" in types, ( + f"Nova Sonic never emitted response.created; types={sorted(types)}" + ) + assert transcript(events).strip() != "" or "response.done" in types, ( + "Nova Sonic response.create produced no transcript (LIT-2239 hang)" + ) + done = parse_last(events, "response.done", ResponseDone) + assert done is not None, ( + f"Nova Sonic never completed response.done within timeout; types={sorted(types)}" + ) diff --git a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py index f882bc5b4e4..13992744f42 100644 --- a/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py +++ b/tests/e2e/llm_translation/test_chat_completions_regression_e2e.py @@ -1,25 +1,36 @@ -"""Live regression net for /chat/completions across the configured providers. +"""Live /chat/completions coverage: the #28991 regression net plus per-provider +OpenAI-compatible translation. GH #28991 broke /chat/completions (and /responses) for most models on some releases: a clean 200 came back but with no real completion. A status check -alone would not have caught it, so each case here asserts the product promise - -a non-empty assistant message and a real model name in the body - across the -three providers wired into the gateway config (OpenAI, Anthropic, Gemini). A -regression that empties the completion for any provider fails that provider's -row here. +alone would not have caught it, so TestChatCompletionsRegression asserts the +product promise - a non-empty assistant message and a real model name in the +body - across the three providers wired into the gateway config (OpenAI, +Anthropic, Gemini). A regression that empties the completion for any provider +fails that provider's row here. + +The per-provider classes below cover the OpenAI-compatible /chat/completions +translation for providers customers reach by registering their own deployment +via /model/new (Cohere, Gemini, hosted_vllm), each deleted on teardown. """ from __future__ import annotations +import os + import pytest -from e2e_config import unique_marker +from e2e_config import require_env, unique_marker from e2e_http import unwrap -from models import ChatBody, ChatMessage +from lifecycle import ResourceManager +from models import ChatBody, ChatMessage, LiteLLMParamsBody from passthrough_client import PassthroughClient pytestmark = pytest.mark.e2e +COHERE_BACKEND = "cohere/command-r-08-2024" +GEMINI_BACKEND = "gemini/gemini-2.5-flash" + CHAT_MODELS: tuple[tuple[str, str], ...] = ( ("gpt-5.5", "openai"), ("claude-haiku-4-5", "anthropic"), @@ -68,3 +79,143 @@ class TestChatCompletionsRegression: assert ( message is not None and message.content and message.content.strip() ), f"{model} ({route}): 200 with an empty completion (#28991): {response}" + + +class TestCohereChat: + """Cohere via the OpenAI-compatible /chat/completions path.""" + + @pytest.mark.covers( + "llm.chat_completions.cohere.basic.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_cohere_chat_returns_content( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + (cohere_key,) = require_env("COHERE_API_KEY") + model = f"e2e-cohere-chat-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody(model=COHERE_BACKEND, api_key=cohere_key), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ + ChatMessage( + role="user", + content=f"Reply with the single word pong. {unique_marker()}", + ) + ], + max_tokens=32, + ), + ) + ) + assert response.choices, f"cohere chat returned no choices: {response}" + content = response.choices[0].message.content if response.choices[0].message else None + assert content and content.strip(), f"cohere empty content: {response}" + + +class TestGeminiChatCompletions: + """Gemini via the OpenAI-compatible /chat/completions path, with cost logging. + + Complements the native /gemini passthrough suite by covering the translation + path customers use when they keep the OpenAI SDK. + """ + + @pytest.mark.covers( + "llm.chat_completions.gemini.basic.nonstream.works", + "llm.chat_completions.gemini.basic.nonstream.cost_logged", + exercised_on=["chat_completions"], + ) + def test_gemini_chat_returns_content_and_logs_cost( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + model = f"e2e-gemini-chat-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody(model=GEMINI_BACKEND, api_key="os.environ/GEMINI_API_KEY"), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + tag = f"e2e-gemini-chat-{unique_marker()}" + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ + ChatMessage( + role="user", + content=f"Reply with the single word pong. marker={tag}", + ) + ], + max_tokens=32, + ), + ) + ) + assert response.choices, f"gemini chat returned no choices: {response}" + content = response.choices[0].message.content if response.choices[0].message else None + assert content, f"gemini chat returned empty content: {response}" + + rows = client.proxy.poll_logs_for_key( + key, + min_rows=1, + predicate=lambda rs: any((r.spend or 0) > 0 for r in rs), + ) + assert rows, f"no SpendLogs row for gemini chat on key ending ...{key[-6:]}" + row = rows[0] + assert (row.spend or 0) > 0, f"gemini chat was not costed: {row}" + assert row.status == "success", f"gemini chat spend status={row.status!r}" + + +class TestHostedVllmChat: + """hosted_vllm (self-hosted OpenAI-compatible server) via /chat/completions.""" + + @pytest.mark.covers( + "llm.chat_completions.hosted_vllm.basic.nonstream.works", + exercised_on=["chat_completions"], + ) + def test_hosted_vllm_chat_returns_content( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + (api_base,) = require_env("HOSTED_VLLM_API_BASE") + api_key = (os.environ.get("HOSTED_VLLM_API_KEY") or "").strip() or None + backend = ( + os.environ.get("HOSTED_VLLM_MODEL") or "meta-llama/Llama-3.2-3B-Instruct" + ).strip() + model = f"e2e-vllm-chat-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody( + model=f"hosted_vllm/{backend}", + api_base=api_base, + api_key=api_key, + ), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = resources.key() + + response = unwrap( + client.proxy.chat( + key, + ChatBody( + model=model, + messages=[ + ChatMessage( + role="user", + content=f"Reply with the single word pong. {unique_marker()}", + ) + ], + max_tokens=32, + ), + ) + ) + assert response.choices, f"hosted_vllm chat returned no choices: {response}" + content = response.choices[0].message.content if response.choices[0].message else None + assert content and content.strip(), f"hosted_vllm empty content: {response}" diff --git a/tests/e2e/llm_translation/test_passthrough_headers_e2e.py b/tests/e2e/llm_translation/test_passthrough_headers_e2e.py new file mode 100644 index 00000000000..d9fe37c79ed --- /dev/null +++ b/tests/e2e/llm_translation/test_passthrough_headers_e2e.py @@ -0,0 +1,150 @@ +"""Live e2e: custom pass-through endpoints inject configured headers and honor +x-pass-* client headers (prefix stripped) on the way to the upstream. + +The upstream is a real public echo service (httpbin.org/anything). Creating the +route via POST /config/pass_through_endpoint, calling it with a virtual key, and +asserting the echo body is the product path operators use; a mock would not +prove the proxy actually rewrote the outbound request. +""" + +from __future__ import annotations + +import pytest +from pydantic import BaseModel, Field, ValidationError + +from e2e_config import unique_marker +from e2e_http import AuthHeaders, NoBody, StreamingResponse, require_successful_call, unwrap +from lifecycle import ResourceManager +from models import KeyGenerateBody +from passthrough_client import PassthroughClient + +pytestmark = pytest.mark.e2e + +ECHO_TARGET = "https://httpbin.org/anything" +STATIC_HEADER_NAME = "x-e2e-static-header" +PASS_HEADER_STEM = "e2e-client-marker" +PASS_HEADER_NAME = f"x-pass-{PASS_HEADER_STEM}" + + +class PassThroughCreateBody(BaseModel): + path: str + target: str + headers: dict[str, str] = {} + auth: bool = True + include_subpath: bool = False + + +class PassThroughEndpoint(BaseModel): + id: str | None = None + path: str + target: str + + +class PassThroughCreateResponse(BaseModel): + endpoints: list[PassThroughEndpoint] + + +class PassThroughDeleteParams(BaseModel): + endpoint_id: str + + +class EchoCallHeaders(AuthHeaders): + content_type: str = Field(default="application/json", serialization_alias="Content-Type") + x_pass_e2e_client_marker: str = Field(serialization_alias="x-pass-e2e-client-marker") + + +class EchoBody(BaseModel): + ping: str + + +class EchoResponse(BaseModel): + headers: dict[str, str] + + +def _create_passthrough( + client: PassthroughClient, *, path: str, static_value: str +) -> PassThroughEndpoint: + created = unwrap( + client.proxy.transport.post( + "/config/pass_through_endpoint", + headers=client.proxy.transport.master, + json=PassThroughCreateBody( + path=path, + target=ECHO_TARGET, + headers={STATIC_HEADER_NAME: static_value}, + ), + response_type=PassThroughCreateResponse, + ) + ) + assert created.endpoints, "create returned no endpoints" + endpoint = created.endpoints[0] + assert endpoint.id, "created pass-through endpoint has no id" + return endpoint + + +def _delete_passthrough(client: PassthroughClient, endpoint_id: str) -> None: + _ = client.proxy.transport.delete( + "/config/pass_through_endpoint", + headers=client.proxy.transport.master, + json=NoBody(), + params=PassThroughDeleteParams(endpoint_id=endpoint_id), + response_type=PassThroughCreateResponse, + ) + + +def _echo_headers(resp: StreamingResponse) -> dict[str, str]: + try: + echo = EchoResponse.model_validate_json(resp.body) + except ValidationError as exc: + pytest.fail(f"echo upstream did not return a headers map: {exc}; body={resp.body[:300]}") + return {k.lower(): v for k, v in echo.headers.items()} + + +class TestPassthroughHeaders: + @pytest.mark.covers( + "other.config.passthrough.headers_forwarded", + exercised_on=[], + ) + def test_static_and_x_pass_headers_reach_upstream( + self, client: PassthroughClient, resources: ResourceManager + ) -> None: + marker = unique_marker() + path = f"/e2e-passthrough-headers-{marker}" + static_value = f"static-{marker}" + client_value = f"client-{marker}" + + endpoint = _create_passthrough(client, path=path, static_value=static_value) + assert endpoint.id is not None + resources.defer(lambda: _delete_passthrough(client, endpoint.id or "")) + + key = client.proxy.generate_key( + KeyGenerateBody( + models=[], + allowed_passthrough_routes=[path], + user_id=f"e2e-pass-headers-{marker}", + ) + ) + resources.defer(lambda: client.proxy.delete_key(key)) + + result = client.proxy.transport.send( + path, + headers=EchoCallHeaders( + authorization=f"Bearer {key}", + x_pass_e2e_client_marker=client_value, + ), + json=EchoBody(ping=marker), + ) + require_successful_call(result) + + upstream = _echo_headers(result) + assert upstream.get(STATIC_HEADER_NAME) == static_value, ( + f"configured pass-through header {STATIC_HEADER_NAME!r} not on upstream " + f"request; got {upstream}" + ) + assert upstream.get(PASS_HEADER_STEM) == client_value, ( + f"x-pass-* header should strip the prefix and forward as {PASS_HEADER_STEM!r}; " + f"got {upstream}" + ) + assert PASS_HEADER_NAME not in upstream, ( + "upstream must not see the x-pass- prefix; proxy should strip it" + ) diff --git a/tests/e2e/llm_translation/test_responses_metadata_e2e.py b/tests/e2e/llm_translation/test_responses_metadata_e2e.py new file mode 100644 index 00000000000..6cf24348095 --- /dev/null +++ b/tests/e2e/llm_translation/test_responses_metadata_e2e.py @@ -0,0 +1,122 @@ +"""Live e2e: /v1/responses with store + metadata (LIT-1201 customer path). + +Customers attach metadata and store=true, then continue with previous_response_id. +Both turns must succeed, and any Redis keys written for the session must carry a +positive TTL (not unbounded). +""" + +from __future__ import annotations + +import os +import socket +import time + +import pytest +from pydantic import BaseModel, ConfigDict + +from e2e_config import require_env, unique_marker +from e2e_http import require_successful_call +from endpoints_client import EndpointsClient, ResponsesResult +from lifecycle import ResourceManager +from models import LiteLLMParamsBody + +pytestmark = pytest.mark.e2e + + +class ResponsesMetadataBody(BaseModel): + model: str + input: str + store: bool = True + metadata: dict[str, str] + previous_response_id: str | None = None + instructions: str | None = "You are a helpful assistant." + + +class RedisKeyInfo(BaseModel): + model_config = ConfigDict(frozen=True) + + key: str + ttl: int + + +def _redis_scan(marker: str) -> tuple[RedisKeyInfo, ...]: + import redis + + (host,) = require_env("REDIS_HOST") + port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379") + try: + with socket.create_connection((host, port), timeout=3): + pass + except OSError as exc: + raise AssertionError( + f"REDIS_HOST={host!r}:{port} unreachable ({exc}); " + "LIT-1201 TTL check needs Redis the proxy writes to." + ) from exc + + client = redis.Redis(host=host, port=port, decode_responses=True, socket_timeout=5) + found: list[RedisKeyInfo] = [] + for key in client.scan_iter(match=f"*{marker}*", count=200): + found.append(RedisKeyInfo(key=str(key), ttl=int(client.ttl(key)))) + return tuple(found) + + +class TestResponsesMetadata: + @pytest.mark.covers( + "llm.responses.openai.basic.nonstream.works", + "other.config.responses.metadata_redis_ttl_bounded", + exercised_on=["responses"], + ) + def test_store_metadata_continues_and_redis_keys_have_ttl( + self, endpoints_client: EndpointsClient, resources: ResourceManager + ) -> None: + # Anthropic avoids OpenAI/Gemini quota flakes; Responses translation still + # exercises store + metadata + previous_response_id on the proxy. + marker = unique_marker() + model = f"e2e-resp-meta-{marker}" + model_id = endpoints_client.create_model( + model, + LiteLLMParamsBody( + model="anthropic/claude-haiku-4-5-20251001", + api_key="os.environ/ANTHROPIC_API_KEY", + ), + ) + resources.defer(lambda: endpoints_client.delete_model(model_id)) + key = resources.key() + + first = endpoints_client.proxy.transport.send( + "/v1/responses", + headers=endpoints_client.proxy.transport.bearer(key), + json=ResponsesMetadataBody( + model=model, + input=f"Remember marker {marker}. Reply with one word.", + metadata={"session_id": marker, "customer": "e2e"}, + ), + ) + require_successful_call(first) + parsed = ResponsesResult.model_validate_json(first.body) + assert parsed.id, f"responses must return an id: {first.body[:300]}" + assert parsed.text.strip(), f"responses returned empty text: {first.body[:300]}" + + second = endpoints_client.proxy.transport.send( + "/v1/responses", + headers=endpoints_client.proxy.transport.bearer(key), + json=ResponsesMetadataBody( + model=model, + input="Reply with the single word ok.", + previous_response_id=parsed.id, + metadata={"session_id": marker, "turn": "2"}, + ), + ) + require_successful_call(second) + second_parsed = ResponsesResult.model_validate_json(second.body) + assert second_parsed.text.strip(), ( + f"previous_response_id follow-up returned empty text: {second.body[:300]}" + ) + + time.sleep(1.0) + keys = _redis_scan(marker) + unbounded = tuple(k for k in keys if k.ttl == -1) + assert not unbounded, ( + "responses metadata must not leave Redis keys without TTL (LIT-1201); " + f"unbounded={unbounded}" + ) diff --git a/tests/e2e/logging/conftest.py b/tests/e2e/logging/conftest.py index 2285eb8d695..60536ea01d4 100644 --- a/tests/e2e/logging/conftest.py +++ b/tests/e2e/logging/conftest.py @@ -10,7 +10,7 @@ import os import pytest -from logging_client import LangfuseCreds, LoggingClient, build_logging_client, load_langfuse_creds +from logging_client import LoggingClient, build_logging_client from datadog_reader import DdLogsReader, build_dd_logs_reader from otel_client import OtelReader, build_otel_reader from proxy_client import ProxyClient @@ -19,15 +19,14 @@ from proxy_client import ProxyClient def pytest_configure(config: pytest.Config) -> None: config.addinivalue_line( "markers", - "covers: registry cell a test covers, e.g. logging.langfuse.success.logs_spend", + "covers: registry cell a test covers, e.g. logging.datadog.success.exports_metric", ) @pytest.fixture(scope="session") def client(proxy: ProxyClient) -> LoggingClient: """The logging suite's client: holds the shared ProxyClient so `resources` / - `scoped_key` clean up keys and teams, and adds `/metrics` scraping plus - Langfuse read-back.""" + `scoped_key` clean up keys and teams, and adds `/metrics` scraping.""" return build_logging_client(proxy) @@ -51,9 +50,3 @@ def datadog_creds() -> None: pytest.fail( "Datadog e2e requires DD_API_KEY and DD_SITE; missing credentials is a hard failure, not a skip" ) - - -@pytest.fixture(scope="session") -def langfuse_creds() -> LangfuseCreds: - """Require real Langfuse cloud credentials for team callback + trace poll.""" - return load_langfuse_creds() diff --git a/tests/e2e/management/management_client.py b/tests/e2e/management/management_client.py index e967fb7b504..a9dedac8e61 100644 --- a/tests/e2e/management/management_client.py +++ b/tests/e2e/management/management_client.py @@ -14,32 +14,45 @@ from e2e_http import NoBody, ProbeResult, Result, StreamingResponse, Success, Un from models import ( ChatBody, ChatMessage, + KeyBlockBody, KeyDeleteBody, KeyGenerateBody, + KeyGenerateResponse, KeyListParams, KeyListResponse, + KeyRegenerateBody, KeyUpdateBody, + ModelDeleteBody, OrgDeleteBody, OrgInfoParams, OrgInfoResponse, OrgNewBody, OrgNewResponse, + OrgUpdateBody, + TagDeleteBody, + TagListEntry, + TagListResponse, + TagNewBody, TeamData, TeamDeleteBody, TeamInfoParams, TeamInfoResponse, + TeamListResponse, TeamMemberAddBody, TeamMemberDeleteBody, TeamMemberEntry, TeamNewBody, TeamNewResponse, + TeamUpdateBody, UserDeleteBody, + UserDeleteResponse, UserInfoParams, UserInfoResponse, UserListParams, UserListResponse, UserNewBody, UserNewResponse, + UserUpdateBody, ) MODEL_ACCESS_DENIED_MARKER = "key_model_access_denied" @@ -89,6 +102,37 @@ class ManagementClient: ) ) + def delete_model_strict(self, model_id: str) -> None: + """Strict delete for the act phase of a test: a failed delete is a hard + failure, unlike the warn-only ProxyClient.delete_model used at teardown.""" + _ = unwrap( + self.proxy.transport.post( + "/model/delete", + headers=self.proxy.transport.master, + json=ModelDeleteBody(id=model_id), + response_type=NoBody, + ) + ) + + def block_key(self, key: str) -> None: + _ = unwrap( + self.proxy.transport.post( + "/key/block", + headers=self.proxy.transport.master, + json=KeyBlockBody(key=key), + response_type=NoBody, + ) + ) + def regenerate_key(self, key: str) -> str: + return unwrap( + self.proxy.transport.post( + "/key/regenerate", + headers=self.proxy.transport.master, + json=KeyRegenerateBody(key=key), + response_type=KeyGenerateResponse, + ) + ).key + def key_alias_count(self, key_alias: str) -> int: return unwrap( self.proxy.transport.get( @@ -111,6 +155,28 @@ class ManagementClient: self._wait_for_team(team_id) return team_id + def update_team(self, body: TeamUpdateBody) -> None: + last: Result[NoBody] | None = None + for attempt in range(5): + last = self.proxy.transport.post( + "/team/update", + headers=self.proxy.transport.master, + json=body, + response_type=NoBody, + ) + match last: + case Success(): + return + case UnknownApiError(body=body_text) if ( + "connecting to redis" in body_text.lower() or "name resolution" in body_text.lower() + ): + time.sleep(0.5 * (attempt + 1)) + continue + case _: + break + assert last is not None + raise AssertionError(last) + def delete_team(self, team_id: str) -> None: _ = self.proxy.transport.post( "/team/delete", @@ -129,6 +195,19 @@ class ManagementClient: ) ).team_info + def team_list_ids(self) -> tuple[str, ...]: + return tuple( + entry.team_id + for entry in unwrap( + self.proxy.transport.get( + "/team/list", + headers=self.proxy.transport.master, + params=NoBody(), + response_type=TeamListResponse, + ) + ).root + ) + def team_info_status(self, team_id: str) -> ProbeResult: return self.proxy.transport.probe("/team/info", params=TeamInfoParams(team_id=team_id)) @@ -191,6 +270,16 @@ class ManagementClient: ) ).user_id + def update_user(self, body: UserUpdateBody) -> None: + _ = unwrap( + self.proxy.transport.post( + "/user/update", + headers=self.proxy.transport.master, + json=body, + response_type=NoBody, + ) + ) + def delete_user(self, user_id: str) -> None: _ = self.proxy.transport.post( "/user/delete", @@ -199,6 +288,18 @@ class ManagementClient: response_type=NoBody, ) + def delete_user_strict(self, user_id: str) -> None: + """Strict delete for the act phase of a test: a failed delete is a hard + failure, unlike the warn-only delete_user used at teardown.""" + _ = unwrap( + self.proxy.transport.post( + "/user/delete", + headers=self.proxy.transport.master, + json=UserDeleteBody(user_ids=[user_id]), + response_type=UserDeleteResponse, + ) + ) + def user_info(self, user_id: str) -> UserInfoResponse: return unwrap( self.proxy.transport.get( @@ -219,6 +320,17 @@ class ManagementClient: ) ).total + def user_list_ids(self, user_id: str) -> tuple[str, ...]: + listing = unwrap( + self.proxy.transport.get( + "/user/list", + headers=self.proxy.transport.master, + params=UserListParams(user_ids=user_id), + response_type=UserListResponse, + ) + ) + return tuple(row.user_id for row in listing.users) + def create_org(self, body: OrgNewBody) -> str: return unwrap( self.proxy.transport.post( @@ -229,6 +341,16 @@ class ManagementClient: ) ).organization_id + def update_org(self, body: OrgUpdateBody) -> None: + _ = unwrap( + self.proxy.transport.patch( + "/organization/update", + headers=self.proxy.transport.master, + json=body, + response_type=NoBody, + ) + ) + def delete_org(self, organization_id: str) -> None: _ = self.proxy.transport.delete( "/organization/delete", @@ -247,6 +369,38 @@ class ManagementClient: ) ) + def org_info_status(self, organization_id: str) -> ProbeResult: + return self.proxy.transport.probe("/organization/info", params=OrgInfoParams(organization_id=organization_id)) + def create_tag(self, body: TagNewBody) -> None: + _ = unwrap( + self.proxy.transport.post( + "/tag/new", + headers=self.proxy.transport.master, + json=body, + response_type=NoBody, + ) + ) + + def delete_tag(self, name: str) -> None: + _ = self.proxy.transport.post( + "/tag/delete", + headers=self.proxy.transport.master, + json=TagDeleteBody(name=name), + response_type=NoBody, + ) + + def tag_list(self) -> tuple[TagListEntry, ...]: + return tuple( + unwrap( + self.proxy.transport.get( + "/tag/list", + headers=self.proxy.transport.master, + params=NoBody(), + response_type=TagListResponse, + ) + ).root + ) + def chat_status(self, key: str, model: str, content: str) -> StreamingResponse: return self.proxy.transport.send( "/chat/completions", diff --git a/tests/e2e/management/test_management_e2e.py b/tests/e2e/management/test_management_e2e.py index adbf3e8b065..18bc384a879 100644 --- a/tests/e2e/management/test_management_e2e.py +++ b/tests/e2e/management/test_management_e2e.py @@ -9,6 +9,7 @@ so the traffic-facing read-backs poll to a deadline instead of asserting once. from __future__ import annotations +import math import time from collections.abc import Callable @@ -22,7 +23,7 @@ from management_client import ( ROUTE_NOT_ALLOWED_MARKER, ManagementClient, ) -from models import KeyGenerateBody, OrgNewBody, TeamNewBody, UserNewBody +from models import KeyGenerateBody, OrgInfoResponse, OrgNewBody, OrgUpdateBody, TagListEntry, TagNewBody, TeamNewBody, TeamUpdateBody, UserNewBody, UserUpdateBody, LiteLLMParamsBody, ModelInfoEntry pytestmark = pytest.mark.e2e @@ -168,6 +169,61 @@ class TestKeyRoutes: _ = _poll(client, rejected, "deleted key was still accepted on chat (never rejected 401) at the deadline") + @pytest.mark.covers("mgmt.key.list.happy_path") + def test_created_key_appears_in_key_list_inventory( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + alias = f"e2e-mgmt-keylist-{unique_marker()}" + assert client.key_alias_count(alias) == 0, ( + f"/key/list already reports a key under the unused alias {alias!r} before it is created" + ) + + _ = _generate_key(client, resources, KeyGenerateBody(key_alias=alias)) + + def listed() -> bool | None: + return True if client.key_alias_count(alias) == 1 else None + + _ = _poll( + client, listed, f"created key with alias {alias!r} never appeared in /key/list before the deadline" + ) + + + @pytest.mark.covers("mgmt.key.block.persists") + def test_block_persists_to_key_info(self, client: ManagementClient, resources: ResourceManager) -> None: + key = _generate_key(client, resources, KeyGenerateBody(models=["gemini-2.5-flash"])) + assert not client.proxy.key_info(key).blocked, "/key/info reports the key blocked before /key/block ran" + + client.block_key(key) + + def blocked() -> bool | None: + return True if client.proxy.key_info(key).blocked else None + + _ = _poll(client, blocked, "/key/info never reported the key blocked after /key/block before the deadline") +class TestKeyRegeneration: + @pytest.mark.covers("mgmt.key.regenerate.happy_path") + def test_regenerate_rotates_to_a_working_new_key( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + old_key = _generate_key(client, resources, KeyGenerateBody(models=["gpt-5.5"])) + + new_key = client.regenerate_key(old_key) + resources.defer(lambda: client.proxy.delete_key(new_key)) + assert new_key != old_key, "regenerate returned the same key string, so no rotation happened" + + def new_accepted() -> bool | None: + outcome = client.chat_status(new_key, "gpt-5.5", f"say hi {unique_marker()}") + return True if outcome.status_code != 401 else None + + _ = _poll(client, new_accepted, "regenerated key was never accepted at auth (still 401) at the deadline") + + def old_rejected() -> bool | None: + outcome = client.chat_status(old_key, "gpt-5.5", f"say hi {unique_marker()}") + return True if outcome.status_code == 401 else None + + _ = _poll( + client, old_rejected, "old key was still accepted after regeneration (never rejected 401) at the deadline" + ) + class TestTeamRoutes: @pytest.mark.covers("mgmt.team.new.persists") @@ -189,6 +245,61 @@ class TestTeamRoutes: f"key generated under team {team_id} carries team_id {key_info.team_id!r} in /key/info" ) + @pytest.mark.covers("mgmt.team.update.persists") + def test_update_persists_to_team_info(self, client: ManagementClient, resources: ResourceManager) -> None: + team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", ["gemini-2.5-flash"]) + + updated_alias = f"e2e-mgmt-team-updated-{unique_marker()}" + client.update_team(TeamUpdateBody(team_id=team_id, team_alias=updated_alias)) + + def reflected() -> bool | None: + return True if client.team_info(team_id).team_alias == updated_alias else None + + _ = _poll(client, reflected, f"/team/info never reflected team_alias {updated_alias!r} after /team/update") + @pytest.mark.covers("mgmt.team.list.happy_path") + def test_created_team_appears_in_team_list( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + alias = f"e2e-mgmt-team-{unique_marker()}" + team_id = _create_team(client, resources, alias, ["gemini-2.5-flash"]) + + _ = _poll( + client, + lambda: team_id if team_id in client.team_list_ids() else None, + f"/team/list never included the created team {team_id}", + ) + + @pytest.mark.covers("mgmt.team.delete.persists") + def test_delete_persists_and_revokes_team_bound_key( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """The teardown's deferred delete_team/delete_key fire again on the already- + deleted team and key by design: both are warn-only no-ops, and the deferred + cleanup must survive this test failing before the in-body delete.""" + team_id = _create_team(client, resources, f"e2e-mgmt-team-{unique_marker()}", ["gpt-5.5"]) + key = _generate_key(client, resources, KeyGenerateBody(team_id=team_id)) + + def accepted() -> bool | None: + outcome = client.chat_status(key, "gpt-5.5", f"say hi {unique_marker()}") + return True if outcome.status_code != 401 else None + + _ = _poll(client, accepted, "team-bound key was never accepted at auth before team deletion") + + client.delete_team(team_id) + + probe = client.team_info_status(team_id) + assert probe.status_code == 404, ( + f"deleted team {team_id} still resolves: /team/info returned {probe.status_code}: {probe.body[:300]}" + ) + + def rejected() -> bool | None: + outcome = client.chat_status(key, "gpt-5.5", f"say hi {unique_marker()}") + return True if outcome.status_code == 401 else None + + _ = _poll( + client, rejected, "team-bound key was still accepted on chat (never rejected 401) after team deletion" + ) + @pytest.mark.covers("mgmt.team.member_add.persists") def test_member_add_and_delete_persist_to_team_info( self, client: ManagementClient, resources: ResourceManager @@ -226,6 +337,64 @@ class TestUserRoutes: f"/user/info reports user_role {info.user_role!r}, configured 'internal_user'" ) + @pytest.mark.covers("mgmt.user.update.persists") + def test_update_persists_to_user_info(self, client: ManagementClient, resources: ResourceManager) -> None: + email = f"e2e-mgmt-{unique_marker()}@example.com" + user_id = _create_user(client, resources, UserNewBody(user_email=email, user_role="internal_user")) + + before = client.user_info(user_id).user_info + assert before.user_role == "internal_user", ( + f"/user/info reports pre-update user_role {before.user_role!r}, expected 'internal_user'" + ) + + client.update_user(UserUpdateBody(user_id=user_id, user_role="internal_user_viewer")) + + info = client.user_info(user_id).user_info + assert info.user_role == "internal_user_viewer", ( + f"/user/info reports user_role {info.user_role!r} after /user/update to 'internal_user_viewer'" + ) + @pytest.mark.covers("mgmt.user.delete.persists") + def test_delete_removes_the_user_from_inventory( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """The teardown's deferred delete fires again on the already-deleted user by + design: the deferred cleanup must survive this test failing before the + in-body delete, and a repeat /user/delete is a cheap no-op the warn-only + teardown absorbs.""" + user_id = _create_user( + client, + resources, + UserNewBody(user_email=f"e2e-mgmt-{unique_marker()}@example.com", user_role="internal_user"), + ) + assert client.user_count(user_id) == 1, f"user {user_id} was not created before deletion" + + client.delete_user_strict(user_id) + + def removed() -> bool | None: + return True if client.user_count(user_id) == 0 else None + + _ = _poll(client, removed, f"user {user_id} still present in /user/list after /user/delete at the deadline") + + @pytest.mark.covers("mgmt.user.list.happy_path") + def test_created_users_appear_in_user_list( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + user_ids = tuple( + _create_user( + client, + resources, + UserNewBody(user_email=f"e2e-mgmt-{unique_marker()}@example.com", user_role="internal_user"), + ) + for _ in range(2) + ) + + for user_id in user_ids: + _ = _poll( + client, + lambda user_id=user_id: (True if user_id in client.user_list_ids(user_id) else None), + f"/user/list never listed the created user {user_id} in the admin inventory", + ) + class TestOrganizationRoutes: @pytest.mark.covers("mgmt.organization.new.happy_path") @@ -244,6 +413,160 @@ class TestOrganizationRoutes: f"/organization/info reports models {info.models}, configured ['gemini-2.5-flash']" ) + @pytest.mark.covers("mgmt.organization.update.persists") + def test_update_alias_persists_to_organization_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + org_id = client.create_org(OrgNewBody(organization_alias=f"e2e-mgmt-org-{unique_marker()}")) + resources.defer(lambda: client.delete_org(org_id)) + + new_alias = f"e2e-mgmt-org-{unique_marker()}" + client.update_org(OrgUpdateBody(organization_id=org_id, organization_alias=new_alias)) + + def attempt() -> OrgInfoResponse | None: + info = client.org_info(org_id) + return info if info.organization_alias == new_alias else None + + _ = _poll( + client, attempt, f"/organization/info never reflected updated alias {new_alias!r} before the deadline" + ) + + @pytest.mark.covers("mgmt.organization.delete.persists") + def test_delete_removes_from_organization_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """The teardown's deferred delete fires again on the already-deleted org by + design: the deferred cleanup must survive this test failing before the + in-body delete, and a repeat /organization/delete is a warn-only no-op the + teardown absorbs.""" + org_id = client.create_org(OrgNewBody(organization_alias=f"e2e-mgmt-org-{unique_marker()}")) + resources.defer(lambda: client.delete_org(org_id)) + + assert client.org_info_status(org_id).status_code == 200, ( + f"/organization/info did not resolve org {org_id} before deletion" + ) + + client.delete_org(org_id) + + def gone() -> bool | None: + return True if client.org_info_status(org_id).status_code == 404 else None + + _ = _poll(client, gone, f"org {org_id} still resolved on /organization/info after /organization/delete") + + +class TestTagRoutes: + @pytest.mark.covers("mgmt.tag.new.happy_path") + def test_new_persists_to_tag_list(self, client: ManagementClient, resources: ResourceManager) -> None: + name = f"e2e-mgmt-tag-{unique_marker()}" + description = "Tag for spend categorization" + + assert all(entry.name != name for entry in client.tag_list()), ( + f"tag {name!r} was already listed by /tag/list before /tag/new created it" + ) + + client.create_tag(TagNewBody(name=name, description=description)) + resources.defer(lambda: client.delete_tag(name)) + + def listed() -> TagListEntry | None: + return next((entry for entry in client.tag_list() if entry.name == name), None) + + entry = _poll(client, listed, f"/tag/list never listed {name!r} after /tag/new") + assert entry.description == description, ( + f"/tag/list reports description {entry.description!r} for {name!r}, configured {description!r}" + ) + + +_INITIAL_INPUT_COST = 0.00000111 +_UPDATED_INPUT_COST = 0.00000222 + + +def _model_entry(client: ManagementClient, model_name: str) -> ModelInfoEntry | None: + return next((entry for entry in client.proxy.model_info() if entry.model_name == model_name), None) + + +class TestModelRoutes: + @pytest.mark.covers("mgmt.model.update.persists") + def test_update_persists_input_cost_to_model_info( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + model_name = f"e2e-mgmt-model-{unique_marker()}" + model_id = client.proxy.create_model( + model_name, + LiteLLMParamsBody( + model="gpt-4o-mini", + mock_response="ok", + input_cost_per_token=_INITIAL_INPUT_COST, + ), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + before = _model_entry(client, model_name) + assert before is not None, f"{model_name} absent from /model/info right after /model/new" + initial = before.litellm_params.input_cost_per_token + assert initial is not None and math.isclose(initial, _INITIAL_INPUT_COST, rel_tol=1e-9), ( + f"/model/info reports input_cost_per_token {initial}, registered {_INITIAL_INPUT_COST}" + ) + + client.proxy.update_model( + model_id, + LiteLLMParamsBody(model="gpt-4o-mini", input_cost_per_token=_UPDATED_INPUT_COST), + ) + + def updated() -> ModelInfoEntry | None: + entry = _model_entry(client, model_name) + if entry is None: + return None + cost = entry.litellm_params.input_cost_per_token + if cost is not None and math.isclose(cost, _UPDATED_INPUT_COST, rel_tol=1e-9): + return entry + return None + + _ = _poll( + client, + updated, + f"/model/info never reported input_cost_per_token {_UPDATED_INPUT_COST} for {model_name} " + "after /model/update", + ) + + @pytest.mark.covers("mgmt.model.delete.persists") + def test_delete_removes_from_model_info_catalog( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + """The teardown's deferred delete fires again on the already-deleted model by + design: it is the safety net if this test fails before the in-body delete, and + a repeat /model/delete is a warn-only no-op the teardown absorbs.""" + model_name = f"e2e-mgmt-model-{unique_marker()}" + model_id = client.proxy.create_model(model_name, LiteLLMParamsBody(model="openai/gpt-5.5", api_key="dummy")) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + assert model_name in [entry.model_name for entry in client.proxy.model_info()], ( + f"{model_name} absent from /model/info right after /model/new; cannot prove deletion removes it" + ) + + client.delete_model_strict(model_id) + + def absent() -> bool | None: + return True if model_name not in [entry.model_name for entry in client.proxy.model_info()] else None + + _ = _poll(client, absent, f"{model_name} still present in /model/info after /model/delete at the deadline") + + @pytest.mark.covers("mgmt.model.add.persists") + def test_new_persists_to_model_info_catalog( + self, client: ManagementClient, resources: ResourceManager + ) -> None: + model_name = f"e2e-mgmt-model-{unique_marker()}" + model_id = client.proxy.create_model( + model_name, + LiteLLMParamsBody(model="openai/gpt-5.5", api_key="e2e-dummy-key"), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + cataloged = [entry.model_name for entry in client.proxy.model_info()] + assert model_name in cataloged, ( + f"/model/info does not list {model_name!r} after /model/new; registration did not persist " + f"into the routing catalog: {cataloged}" + ) + def _assert_route_forbidden(route: str, outcome: StreamingResponse) -> None: assert outcome.status_code == 403, ( diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 2e7bfe41e30..28cc7984598 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -65,6 +65,7 @@ class KeyGenerateBody(BaseModel): tpm_limit: int | None = None rpm_limit: int | None = None allowed_routes: list[str] | None = None + allowed_passthrough_routes: list[str] | None = None metadata: KeyMetadata | None = None object_permission: ObjectPermission | None = None @@ -73,6 +74,10 @@ class KeyGenerateResponse(BaseModel): key: str +class KeyRegenerateBody(BaseModel): + key: str + + class KeyDeleteBody(BaseModel): keys: list[str] @@ -94,6 +99,7 @@ class KeyInfo(BaseModel): tpm_limit: int | None = None rpm_limit: int | None = None team_id: str | None = None + blocked: bool | None = None spend: float | None = None max_budget: float | None = None budget_reset_at: str | None = None @@ -125,6 +131,21 @@ class ChatMessage(BaseModel): content: str +class CacheControl(BaseModel): + type: str = "ephemeral" + + +class TextBlock(BaseModel): + type: str = "text" + text: str + cache_control: CacheControl | None = None + + +class RichMessage(BaseModel): + role: str + content: list[TextBlock] + + class ThinkingParam(BaseModel): """Extended-thinking control shared by Anthropic and DeepSeek reasoner models. DeepSeek accepts only ``type`` (enabled/disabled) and ignores budget_tokens; @@ -161,6 +182,27 @@ class ChatBody(BaseModel): guardrails: list[str] | None = None +class RouterSettingsOverride(BaseModel): + """Per-request `router_settings_override` in a /chat/completions body: the + reliability knobs (fallbacks by trigger, retry count) the reliability suite + drives per call instead of via static router config. Serialized exclude_none, so + an override sets only the strategies a test exercises. Each fallbacks map is + model_name -> the ordered fallback model_names to try.""" + + fallbacks: list[dict[str, list[str]]] | None = None + context_window_fallbacks: list[dict[str, list[str]]] | None = None + content_policy_fallbacks: list[dict[str, list[str]]] | None = None + num_retries: int | None = None + + +class ReliabilityChatBody(ChatBody): + """A /chat/completions body carrying a per-request router_settings_override. + Composes ChatBody (no attribute repetition) and adds the override; serialized + exclude_none so an absent override never leaks into the request.""" + + router_settings_override: RouterSettingsOverride | None = None + + class OutMessage(BaseModel): content: str | None = None reasoning_content: str | None = None @@ -515,12 +557,16 @@ class LiteLLMParamsBody(BaseModel): s3_access_key_id: str | None = None s3_secret_access_key: str | None = None aws_batch_role_arn: str | None = None + aws_role_name: str | None = None + aws_session_name: str | None = None + aws_external_id: str | None = None input_cost_per_token: float | None = None output_cost_per_token: float | None = None extra_headers: dict[str, str] | None = None use_in_pass_through: bool | None = None complexity_router_config: dict[str, object] | None = None mock_response: str | None = None + timeout: float | None = None ModelMode = Literal["batch", "realtime", "image_generation"] @@ -547,6 +593,17 @@ class ModelNewResponse(BaseModel): model_id: str +class ModelUpdateBody(BaseModel): + """POST /model/update body: the target deployment (`model_info.id`) plus the + `litellm_params` to merge over its stored params. The handler overlays only the + non-null fields, so a body carrying `input_cost_per_token` re-prices the + deployment while leaving its other params intact.""" + + model_config = ConfigDict(protected_namespaces=()) + litellm_params: LiteLLMParamsBody + model_info: ModelInfoBody + + class ModelListEntry(BaseModel): id: str @@ -581,6 +638,10 @@ class KeyUpdateBody(BaseModel): models: list[str] +class KeyBlockBody(BaseModel): + key: str + + class KeyListParams(BaseModel): key_alias: str @@ -594,17 +655,27 @@ class TeamMemberEntry(BaseModel): user_id: str +class TeamMetadata(BaseModel): + disable_global_guardrails: bool | None = None + + class TeamNewBody(BaseModel): team_alias: str models: list[str] = [] team_id: str | None = None organization_id: str | None = None + metadata: TeamMetadata | None = None class TeamNewResponse(BaseModel): team_id: str +class TeamUpdateBody(BaseModel): + team_id: str + team_alias: str + + class TeamInfoParams(BaseModel): team_id: str @@ -634,6 +705,15 @@ class TeamDeleteBody(BaseModel): team_ids: list[str] +class TeamListEntry(BaseModel): + team_id: str + + +class TeamListResponse(RootModel[list[TeamListEntry]]): + """GET /team/list answers with a bare array of team objects (not an object + wrapping them). Only team_id is read; pydantic ignores the rest.""" + + UserRole = Literal["proxy_admin", "proxy_admin_viewer", "internal_user", "internal_user_viewer"] @@ -647,6 +727,11 @@ class UserNewResponse(BaseModel): user_id: str +class UserUpdateBody(BaseModel): + user_id: str + user_role: UserRole + + class UserInfoParams(BaseModel): user_id: str @@ -666,11 +751,20 @@ class UserDeleteBody(BaseModel): user_ids: list[str] +class UserDeleteResponse(RootModel[int]): + pass + + class UserListParams(BaseModel): user_ids: str +class UserListRow(BaseModel): + user_id: str + + class UserListResponse(BaseModel): + users: list[UserListRow] total: int @@ -683,6 +777,11 @@ class OrgNewResponse(BaseModel): organization_id: str +class OrgUpdateBody(BaseModel): + organization_id: str + organization_alias: str + + class OrgInfoParams(BaseModel): organization_id: str @@ -695,3 +794,26 @@ class OrgInfoResponse(BaseModel): class OrgDeleteBody(BaseModel): organization_ids: list[str] + + +# ---------- tags (management) ---------- + + +class TagNewBody(BaseModel): + name: str + description: str | None = None + + +class TagDeleteBody(BaseModel): + name: str + + +class TagListEntry(BaseModel): + name: str + description: str | None = None + + +class TagListResponse(RootModel[list[TagListEntry]]): + """GET /tag/list answers with a bare array of tag configs (the stored tags plus + any dynamically-seen spend tags), not an object wrapping them. Read the rows off + .root.""" diff --git a/tests/e2e/proxy_client.py b/tests/e2e/proxy_client.py index f1039257193..6c6b948e29c 100644 --- a/tests/e2e/proxy_client.py +++ b/tests/e2e/proxy_client.py @@ -54,6 +54,7 @@ from models import ( ModelNewBody, ModelNewResponse, ModelsListResponse, + ModelUpdateBody, OcrBody, OcrResponse, SpendLogRow, @@ -211,6 +212,23 @@ class ProxyClient: f"propagation or STORE_MODEL_IN_DB reload issue){last_error}" ) + def update_model(self, model_id: str, litellm_params: LiteLLMParamsBody) -> None: + """Merge `litellm_params` over the deployment `model_id`'s stored params via + POST /model/update. The proxy overlays only the non-null fields and clears + its model cache, so a later /model/info read reflects the change (eventually, + after the reload).""" + unwrap( + self.transport.post( + "/model/update", + headers=self.transport.master, + json=ModelUpdateBody( + litellm_params=litellm_params, + model_info=ModelInfoBody(id=model_id), + ), + response_type=NoBody, + ) + ) + def delete_model(self, model_id: str) -> None: result = self.transport.post( "/model/delete", diff --git a/tests/e2e/quota_management/ratelimit/test_redis_backed_ratelimit_e2e.py b/tests/e2e/quota_management/ratelimit/test_redis_backed_ratelimit_e2e.py new file mode 100644 index 00000000000..ed6f0ce3b2c --- /dev/null +++ b/tests/e2e/quota_management/ratelimit/test_redis_backed_ratelimit_e2e.py @@ -0,0 +1,76 @@ +"""Live e2e: RPM enforcement on the Redis-backed limiter path customers run. + +Requires REDIS_HOST reachable from this process. A key with rpm_limit=1 must +serve the first chat and 429 the second. +""" + +from __future__ import annotations + +import os +import socket + +import pytest + +from e2e_config import require_env, unique_marker +from e2e_http import require_successful_call +from lifecycle import ResourceManager +from models import KeyGenerateBody, LiteLLMParamsBody +from quota_client import QuotaClient + +pytestmark = pytest.mark.e2e + +BACKEND = "anthropic/claude-haiku-4-5-20251001" + + +def _require_redis_reachable() -> None: + (host,) = require_env("REDIS_HOST") + port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379") + try: + with socket.create_connection((host, port), timeout=3): + return + except OSError as exc: + raise AssertionError( + f"REDIS_HOST={host!r} port={port} is not reachable ({exc}). " + "Redis-backed rate limiting e2e needs a live Redis the proxy shares." + ) from exc + + +class TestRedisBackedRateLimit: + @pytest.mark.covers( + "quota_management.ratelimit.redis_backed.blocks_over_limit", + exercised_on=["chat_completions"], + ) + def test_rpm_limit_one_blocks_second_call( + self, client: QuotaClient, resources: ResourceManager + ) -> None: + _require_redis_reachable() + model = f"e2e-redis-rpm-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody(model=BACKEND, api_key="os.environ/ANTHROPIC_API_KEY"), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + key = client.proxy.generate_key( + KeyGenerateBody( + models=[model], + rpm_limit=1, + key_alias=f"e2e-redis-rpm-{unique_marker()}", + ) + ) + resources.defer(lambda: client.proxy.delete_key(key)) + + info = client.proxy.key_info(key) + assert info.rpm_limit == 1, f"key must echo rpm_limit=1: {info}" + + first = client.chat(key, model, f"ping {unique_marker()}") + require_successful_call(first) + + second = client.chat(key, model, f"pong {unique_marker()}") + assert second.status_code == 429, ( + f"second call over rpm_limit=1 must be 429, got {second.status_code}: " + f"{second.body[:300]}" + ) + assert "rate" in second.body.lower() or "limit" in second.body.lower(), ( + f"429 body should name the rate limit: {second.body[:300]}" + ) diff --git a/tests/e2e/quota_management/ratelimit/test_redis_circuit_breaker_e2e.py b/tests/e2e/quota_management/ratelimit/test_redis_circuit_breaker_e2e.py new file mode 100644 index 00000000000..b509ae000f5 --- /dev/null +++ b/tests/e2e/quota_management/ratelimit/test_redis_circuit_breaker_e2e.py @@ -0,0 +1,90 @@ +"""Live e2e: Redis-backed rate limit path stays responsive (LIT-3523 shape). + +With Redis up, burst past rpm_limit=1, then a fresh key must still complete a +chat in well under REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT. +""" + +from __future__ import annotations + +import os +import socket +import time +from concurrent.futures import ThreadPoolExecutor, as_completed + +import pytest + +from e2e_config import require_env, unique_marker +from e2e_http import require_successful_call +from lifecycle import ResourceManager +from models import KeyGenerateBody, LiteLLMParamsBody +from quota_client import QuotaClient + +pytestmark = pytest.mark.e2e + +BACKEND = "anthropic/claude-haiku-4-5-20251001" +RECOVERY_TIMEOUT = float( + os.environ.get("REDIS_CIRCUIT_BREAKER_RECOVERY_TIMEOUT", "60") or "60" +) + + +def _require_redis() -> None: + (host,) = require_env("REDIS_HOST") + port = int((os.environ.get("REDIS_PORT") or "6379").strip() or "6379") + try: + with socket.create_connection((host, port), timeout=3): + return + except OSError as exc: + raise AssertionError( + f"REDIS_HOST={host!r}:{port} unreachable ({exc}); " + "LIT-3523 e2e needs Redis the proxy shares." + ) from exc + + +class TestRedisCircuitBreakerPath: + @pytest.mark.covers( + "reliability.circuit_breaker.redis.trips_then_recovers", + exercised_on=["chat_completions"], + ) + def test_burst_rate_limit_does_not_freeze_fresh_key( + self, client: QuotaClient, resources: ResourceManager + ) -> None: + _require_redis() + model = f"e2e-cb-model-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody(model=BACKEND, api_key="os.environ/ANTHROPIC_API_KEY"), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + hot_key = client.proxy.generate_key( + KeyGenerateBody( + models=[model], + rpm_limit=1, + key_alias=f"e2e-cb-hot-{unique_marker()}", + ) + ) + resources.defer(lambda: client.proxy.delete_key(hot_key)) + cool_key = client.proxy.generate_key( + KeyGenerateBody(models=[model], key_alias=f"e2e-cb-cool-{unique_marker()}") + ) + resources.defer(lambda: client.proxy.delete_key(cool_key)) + + def _hit() -> int: + return client.chat(hot_key, model, f"burst {unique_marker()}").status_code + + with ThreadPoolExecutor(max_workers=8) as pool: + futures = [pool.submit(_hit) for _ in range(12)] + codes = tuple(f.result() for f in as_completed(futures)) + assert any(code == 429 for code in codes), ( + f"expected some 429 under rpm_limit=1 burst, got {codes}" + ) + + started = time.monotonic() + cool = client.chat(cool_key, model, f"fresh {unique_marker()}") + elapsed = time.monotonic() - started + require_successful_call(cool) + assert elapsed < RECOVERY_TIMEOUT * 0.5, ( + f"fresh key chat took {elapsed:.1f}s after redis rate-limit burst; " + f"customers treat hangs near recovery_timeout={RECOVERY_TIMEOUT}s as " + "LIT-3523 circuit-breaker pain" + ) diff --git a/tests/e2e/quota_management/ratelimit/test_tpm_excludes_cached_tokens_e2e.py b/tests/e2e/quota_management/ratelimit/test_tpm_excludes_cached_tokens_e2e.py new file mode 100644 index 00000000000..b0bc6b3508c --- /dev/null +++ b/tests/e2e/quota_management/ratelimit/test_tpm_excludes_cached_tokens_e2e.py @@ -0,0 +1,162 @@ +"""Live e2e: cached prompt tokens must not burn TPM budget (LIT-1930). + +Customer expectation: after a cacheable prefix is warmed, the remaining TPM +budget decreases by non-cached tokens only. If cached tokens still counted, +remaining would drop by the full prompt size. +""" + +from __future__ import annotations + +import time + +import pytest +from pydantic import BaseModel + +from e2e_config import unique_marker +from e2e_http import require_successful_call, unwrap +from lifecycle import ResourceManager +from models import ( + CacheControl, + ChatResponse, + KeyGenerateBody, + LiteLLMParamsBody, + RichMessage, + TextBlock, + Usage, +) +from quota_client import QuotaClient + +pytestmark = pytest.mark.e2e + +# Anthropic prompt caching (host has ANTHROPIC_API_KEY; Bedrock was "Operation not allowed"). +ANTHROPIC_MODEL = "anthropic/claude-haiku-4-5-20251001" +# High enough that pre-call reservation of a cacheable prefix still clears. +TPM_LIMIT = 100_000 + + +class CacheChatBody(BaseModel): + model: str + messages: list[RichMessage] + max_tokens: int = 16 + cache: dict[str, bool] = {"no-cache": True} + + +def _prefix() -> str: + marker = unique_marker() + body = " ".join(f"TPM cache paragraph {i} run {marker}." for i in range(600)) + return f"{body}\nEnd {marker}." + + +def _cached_tokens(usage: Usage | None) -> int: + if usage is None: + return 0 + if usage.cache_read_input_tokens: + return usage.cache_read_input_tokens + if usage.prompt_tokens_details and usage.prompt_tokens_details.cached_tokens: + return usage.prompt_tokens_details.cached_tokens + return 0 + + +def _chat_raw(client: QuotaClient, key: str, model: str, prefix: str): + body = CacheChatBody( + model=model, + messages=[ + RichMessage( + role="system", + content=[TextBlock(text=prefix, cache_control=CacheControl())], + ), + RichMessage(role="user", content=[TextBlock(text="Reply with one word.")]), + ], + ) + return client.proxy.transport.send( + "/chat/completions", + headers=client.proxy.transport.bearer(key), + json=body, + ) + + +def _chat(client: QuotaClient, key: str, model: str, prefix: str) -> ChatResponse: + body = CacheChatBody( + model=model, + messages=[ + RichMessage( + role="system", + content=[TextBlock(text=prefix, cache_control=CacheControl())], + ), + RichMessage(role="user", content=[TextBlock(text="Reply with one word.")]), + ], + ) + return unwrap( + client.proxy.transport.post( + "/chat/completions", + headers=client.proxy.transport.bearer(key), + json=body, + response_type=ChatResponse, + ) + ) + + +class TestTpmExcludesCachedTokens: + @pytest.mark.covers( + "quota_management.ratelimit.tpm.excludes_cached_tokens", + exercised_on=["chat_completions"], + ) + def test_cache_hit_reduces_tpm_by_non_cached_only( + self, client: QuotaClient, resources: ResourceManager + ) -> None: + model = f"e2e-tpm-cache-{unique_marker()}" + model_id = client.proxy.create_model( + model, + LiteLLMParamsBody( + model=ANTHROPIC_MODEL, api_key="os.environ/ANTHROPIC_API_KEY" + ), + ) + resources.defer(lambda: client.proxy.delete_model(model_id)) + key = client.proxy.generate_key( + KeyGenerateBody(models=[model], tpm_limit=TPM_LIMIT) + ) + resources.defer(lambda: client.proxy.delete_key(key)) + + prefix = _prefix() + first = _chat(client, key, model, prefix) + assert first.choices, f"cache prime returned no choices: {first}" + first_total = (first.usage.total_tokens or 0) if first.usage else 0 + assert first_total > 0, f"prime call must report usage: {first.usage}" + + deadline = time.monotonic() + 45.0 + second_usage: Usage | None = None + remaining_after: str | None = None + while time.monotonic() < deadline: + outcome = _chat_raw(client, key, model, prefix) + require_successful_call(outcome) + parsed = ChatResponse.model_validate_json(outcome.body) + if _cached_tokens(parsed.usage) > 0: + second_usage = parsed.usage + remaining_after = outcome.headers.get( + "x-ratelimit-api_key-remaining-tokens" + ) + break + time.sleep(2.0) + + assert second_usage is not None, "second call never reported cache-read tokens" + cached = _cached_tokens(second_usage) + assert cached > 0 + second_total = second_usage.total_tokens or 0 + assert second_total > cached, ( + f"need total > cached so non-cached slice is measurable: {second_usage}" + ) + + assert remaining_after is not None and remaining_after.isdigit(), ( + f"cache-hit response must expose remaining TPM headers, got {remaining_after!r}" + ) + remaining = int(remaining_after) + # If cached tokens were counted, remaining would be limit - first - second_total. + # With exclusion, remaining is closer to limit - first - (second_total - cached). + counted_full = TPM_LIMIT - first_total - second_total + counted_excluding_cache = TPM_LIMIT - first_total - (second_total - cached) + assert remaining > counted_full, ( + f"remaining TPM {remaining} looks like cached tokens still counted " + f"(would be ~{counted_full} if full second_total={second_total} counted; " + f"expected closer to ~{counted_excluding_cache} after excluding " + f"cache_read={cached}; LIT-1930)" + ) diff --git a/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py b/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py index 0d49869aa91..26860212fa3 100644 --- a/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py +++ b/tests/e2e/quota_management/spend_tracking/spend_e2e_client.py @@ -11,7 +11,6 @@ helpers from one place. from __future__ import annotations -import os import time from collections.abc import Callable from dataclasses import dataclass @@ -50,7 +49,6 @@ from models import ( __all__ = [ "SpendClient", "build_client", - "reset_spend_logs", "unique_marker", "unwrap", "is_ok", @@ -59,23 +57,6 @@ __all__ = [ ] -def reset_spend_logs() -> None: - """Truncate LiteLLM_SpendLogs for a clean slate. No proxy endpoint deletes - spend logs (/global/spend/reset keeps them), so go to the DB directly. Uses - DATABASE_URL (default: the local docker postgres on its mapped host port; note - the in-container `@db` host isn't resolvable from the host, so default to - localhost). - """ - import psycopg - - url = os.environ.get( - "DATABASE_URL", - "postgresql://llmproxy:dbpassword9090@localhost:5432/litellm", - ) - with psycopg.connect(url) as conn: - _ = conn.execute('TRUNCATE TABLE "LiteLLM_SpendLogs"') - - def _chat_body( model: str, content: str, diff --git a/tests/e2e/router/conftest.py b/tests/e2e/router/conftest.py index 8ddc19aa94f..98501f9bd7c 100644 --- a/tests/e2e/router/conftest.py +++ b/tests/e2e/router/conftest.py @@ -79,8 +79,8 @@ def _router_is_callable(proxy: ProxyClient) -> bool: return isinstance(result, Success) -@pytest.fixture(scope="session", autouse=True) -def _ensure_complexity_smart_router( # pyright: ignore[reportUnusedFunction] # pytest autouse session fixture, wired by name +@pytest.fixture(scope="session") +def _ensure_complexity_smart_router( # pyright: ignore[reportUnusedFunction] # requested by the complexity test via usefixtures, wired by name client: ComplexityRouterClient, ) -> Iterator[None]: """Ensure the complexity router virtual model exists for this session. diff --git a/tests/e2e/router/reliability_support.py b/tests/e2e/router/reliability_support.py new file mode 100644 index 00000000000..4dab0aaa3fa --- /dev/null +++ b/tests/e2e/router/reliability_support.py @@ -0,0 +1,77 @@ +"""Shared helpers for the reliability e2e tests (fallbacks, timeouts, cache). + +These are plain functions over the router suite's shared ProxyClient, not a +fixture/client class: the tests reuse the router `client` fixture and pass +`client.proxy`. Fallbacks and timeouts are driven by REAL deployments that all +point at the real `openai/gpt-5.5`; a bad base URL yields a real connection +error and a 1ms deadline yields a real timeout, and each test wires the +reroute per request through a `router_settings_override` in the /chat/completions +body, so a single long-lived proxy serves every reliability behavior. +""" + +from __future__ import annotations + +from pydantic import ValidationError + +from proxy_client import ProxyClient +from e2e_http import StreamingResponse +from models import ( + ChatMessage, + ChatResponse, + LiteLLMParamsBody, + ReliabilityChatBody, + RouterSettingsOverride, +) + +REAL_MODEL = "openai/gpt-5.5" +REAL_KEY = "os.environ/OPENAI_API_KEY" + + +def create_bad_base_deployment(proxy: ProxyClient, name: str) -> str: + """Register a deployment pointing at an unreachable base, so every call to it + fails with a real connection error the fallback can reroute around.""" + return proxy.create_model( + name, LiteLLMParamsBody(model=REAL_MODEL, api_key=REAL_KEY, api_base="http://127.0.0.1:9/v1") + ) + + +def create_timeout_deployment(proxy: ProxyClient, name: str) -> str: + """Register a deployment with a 1ms deadline the real backend always exceeds.""" + return proxy.create_model(name, LiteLLMParamsBody(model=REAL_MODEL, api_key=REAL_KEY, timeout=0.001)) + + +def chat_override( + proxy: ProxyClient, + key: str, + model: str, + content: str, + override: RouterSettingsOverride | None = None, + stream: bool = False, +) -> StreamingResponse: + """POST /chat/completions with an optional per-request router_settings_override, + returning the raw outcome so tests read status, body, and reliability headers.""" + return proxy.transport.send( + "/chat/completions", + headers=proxy.transport.bearer(key), + json=ReliabilityChatBody( + model=model, + messages=[ChatMessage(role="user", content=content)], + max_tokens=16, + stream=stream, + router_settings_override=override, + ), + stream=stream, + ) + + +def content_of(resp: StreamingResponse) -> str | None: + """The assistant message content of a successful chat response, or None when the + body is not a success shape (an error body, or an elided streamed body).""" + try: + parsed = ChatResponse.model_validate_json(resp.body) + except ValidationError: + return None + if not parsed.choices: + return None + message = parsed.choices[0].message + return message.content if message is not None else None diff --git a/tests/e2e/router/test_complexity_router_e2e.py b/tests/e2e/router/test_complexity_router_e2e.py index a495e2fdf4d..e8508c963b8 100644 --- a/tests/e2e/router/test_complexity_router_e2e.py +++ b/tests/e2e/router/test_complexity_router_e2e.py @@ -38,6 +38,7 @@ HEURISTIC_TIER_MODELS = frozenset({"openai/gpt-5.5", "gpt-5.5"}) LLM_TIER_MODELS = frozenset({"anthropic/claude-haiku-4-5", "claude-haiku-4-5"}) +@pytest.mark.usefixtures("_ensure_complexity_smart_router") class TestComplexityRouterLlmClassifier: @pytest.mark.skip( reason="product bug LIT-4521: LLM classifier returns SIMPLE for short hard prompts " diff --git a/tests/e2e/router/test_reliability_cache_e2e.py b/tests/e2e/router/test_reliability_cache_e2e.py new file mode 100644 index 00000000000..78d8fcdc08f --- /dev/null +++ b/tests/e2e/router/test_reliability_cache_e2e.py @@ -0,0 +1,37 @@ +"""Live e2e: the response cache returns a cached answer on an exact repeat. + +The same unique prompt is sent twice to the real `gpt-5.5` deployment under the +same key: the first call is a cache miss (the proxy computes and stores the entry, +and returns no x-litellm-cache-key), the second is an exact hit (the proxy serves +from cache and returns x-litellm-cache-key). This relies on the standard Redis +response cache being enabled on the proxy under test. +""" + +from __future__ import annotations + +import pytest + +from complexity_router_client import ComplexityRouterClient +from e2e_config import unique_marker +from reliability_support import chat_override + +pytestmark = pytest.mark.e2e + + +class TestReliabilityCache: + @pytest.mark.covers("reliability.cache.exact.returns_cached") + def test_exact_cache_returns_cached(self, client: ComplexityRouterClient, scoped_key: str) -> None: + prompt = f"cache probe {unique_marker()}" + + first = chat_override(client.proxy, scoped_key, "gpt-5.5", prompt) + assert first.status_code == 200, f"first call should succeed, got {first.status_code}: {first.body[:300]}" + assert "x-litellm-cache-key" not in first.headers, ( + "first (uncached) call must not report a cache-key header" + ) + + second = chat_override(client.proxy, scoped_key, "gpt-5.5", prompt) + assert second.status_code == 200, f"second call should succeed, got {second.status_code}: {second.body[:300]}" + assert "x-litellm-cache-key" in second.headers, ( + "second identical call should hit the response cache and report a cache-key header " + "(requires the proxy's Redis response cache to be enabled)" + ) diff --git a/tests/e2e/router/test_reliability_fallbacks_e2e.py b/tests/e2e/router/test_reliability_fallbacks_e2e.py new file mode 100644 index 00000000000..5b7d21c6ef7 --- /dev/null +++ b/tests/e2e/router/test_reliability_fallbacks_e2e.py @@ -0,0 +1,69 @@ +"""Live e2e: per-request fallbacks reroute a failing deployment's traffic to a +healthy one. + +Each test registers a primary deployment that fails (an unreachable base URL, or +a 1ms deadline) and calls it with a `router_settings_override` mapping it to the +real `gpt-5.5`. The proof the fallback fired is twofold: the response is a real +completion from `gpt-5.5` (a non-empty content string), and the proxy reports at +least one attempted fallback in the x-litellm-attempted-fallbacks header. +""" + +from __future__ import annotations + +import pytest + +from complexity_router_client import ComplexityRouterClient +from e2e_config import unique_marker +from e2e_http import StreamingResponse +from lifecycle import ResourceManager +from models import RouterSettingsOverride +from reliability_support import ( + chat_override, + content_of, + create_bad_base_deployment, + create_timeout_deployment, +) + +pytestmark = pytest.mark.e2e + + +def _assert_served_by_fallback(resp: StreamingResponse) -> None: + assert resp.status_code == 200, f"expected 200 after fallback, got {resp.status_code}: {resp.body[:300]}" + content = content_of(resp) + assert isinstance(content, str) and content, ( + f"the gpt-5.5 fallback should have returned a real completion, got content {content!r} " + f"(body={resp.body[:300]})" + ) + attempted = resp.headers.get("x-litellm-attempted-fallbacks") + assert attempted is not None, "response is missing the x-litellm-attempted-fallbacks header" + assert int(attempted) >= 1, f"x-litellm-attempted-fallbacks should be >= 1, got {attempted!r}" + + +class TestReliabilityFallbacks: + @pytest.mark.covers("reliability.fallback.5xx.routes_to_fallback") + def test_5xx_routes_to_fallback( + self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str + ) -> None: + primary = f"reliability-fail-{unique_marker()}" + model_id = create_bad_base_deployment(client.proxy, primary) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + resp = chat_override( + client.proxy, scoped_key, primary, "say hi", + override=RouterSettingsOverride(fallbacks=[{primary: ["gpt-5.5"]}]), + ) + _assert_served_by_fallback(resp) + + @pytest.mark.covers("reliability.fallback.timeout.routes_to_fallback") + def test_timeout_routes_to_fallback( + self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str + ) -> None: + primary = f"reliability-tofail-{unique_marker()}" + model_id = create_timeout_deployment(client.proxy, primary) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + resp = chat_override( + client.proxy, scoped_key, primary, "say hi", + override=RouterSettingsOverride(fallbacks=[{primary: ["gpt-5.5"]}]), + ) + _assert_served_by_fallback(resp) diff --git a/tests/e2e/router/test_reliability_timeouts_e2e.py b/tests/e2e/router/test_reliability_timeouts_e2e.py new file mode 100644 index 00000000000..f24d5139e66 --- /dev/null +++ b/tests/e2e/router/test_reliability_timeouts_e2e.py @@ -0,0 +1,53 @@ +"""Live e2e: a per-request timeout surfaces to the caller instead of hanging. + +A deployment created with a 1ms deadline always exceeds it against the real +backend. With no fallback in play, the proxy must return the timeout to the +caller: a 408 for a non-streamed request, and the same timeout surfaced on the +streamed path (either a 408 before the stream opens or a timeout error carried in +the response). +""" + +from __future__ import annotations + +import pytest + +from complexity_router_client import ComplexityRouterClient +from e2e_config import unique_marker +from lifecycle import ResourceManager +from reliability_support import chat_override, create_timeout_deployment + +pytestmark = pytest.mark.e2e + + +class TestReliabilityTimeouts: + @pytest.mark.covers("reliability.timeout.request_timeout.exceeds_deadline") + def test_request_timeout_exceeds_deadline( + self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str + ) -> None: + name = f"reliability-timeout-{unique_marker()}" + model_id = create_timeout_deployment(client.proxy, name) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + resp = chat_override(client.proxy, scoped_key, name, "hello") + assert resp.status_code == 408, ( + f"a timed-out request should return 408, got {resp.status_code}: {resp.body[:300]}" + ) + assert "timeout" in resp.body.lower(), f"the 408 body should name the timeout, got: {resp.body[:300]}" + + @pytest.mark.covers("reliability.timeout.stream_timeout.exceeds_deadline") + def test_stream_timeout_exceeds_deadline( + self, client: ComplexityRouterClient, resources: ResourceManager, scoped_key: str + ) -> None: + name = f"reliability-stream-timeout-{unique_marker()}" + model_id = create_timeout_deployment(client.proxy, name) + resources.defer(lambda: client.proxy.delete_model(model_id)) + + resp = chat_override(client.proxy, scoped_key, name, "hello", stream=True) + surfaced = f"{resp.body} {resp.stream_error or ''}".lower() + assert resp.status_code >= 400, ( + f"a timed-out streaming request should surface an error status, got {resp.status_code}: {resp.body[:300]}" + ) + assert "timeout" in surfaced, ( + f"the streamed timeout error should name the timeout, got body={resp.body[:300]}, " + f"stream_error={resp.stream_error!r}" + ) diff --git a/tests/e2e/transport.py b/tests/e2e/transport.py index 10e090f07a9..005b49272e8 100644 --- a/tests/e2e/transport.py +++ b/tests/e2e/transport.py @@ -52,6 +52,16 @@ class Transport(Protocol): ) -> Result[R]: ... def delete[R: BaseModel]( + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + response_type: type[R], + params: BaseModel | None = None, + ) -> Result[R]: ... + + def patch[R: BaseModel]( self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] ) -> Result[R]: ... @@ -121,9 +131,27 @@ class HttpTransport: ) def delete[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + response_type: type[R], + params: BaseModel | None = None, ) -> Result[R]: return e2e_http.delete( + self._url(path), + headers=headers, + json=json, + params=params, + response_type=response_type, + timeout=self.request_timeout, + ) + + def patch[R: BaseModel]( + self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + ) -> Result[R]: + return e2e_http.patch( self._url(path), headers=headers, json=json, @@ -208,6 +236,8 @@ CONTROL_PLANE_PREFIXES: tuple[str, ...] = ( "/model/", "/spend", "/global", + "/config", + "/guardrails", "/openapi.json", ) @@ -265,9 +295,26 @@ class SplitTransport: ) def delete[R: BaseModel]( - self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + self, + path: str, + *, + headers: BaseModel, + json: BaseModel, + response_type: type[R], + params: BaseModel | None = None, ) -> Result[R]: return self._route(path).delete( + path, + headers=headers, + json=json, + response_type=response_type, + params=params, + ) + + def patch[R: BaseModel]( + self, path: str, *, headers: BaseModel, json: BaseModel, response_type: type[R] + ) -> Result[R]: + return self._route(path).patch( path, headers=headers, json=json, response_type=response_type ) diff --git a/tests/mcp_tests/test_aresponses_api_with_mcp.py b/tests/mcp_tests/test_aresponses_api_with_mcp.py index 9cd45f3d6fc..32295310005 100644 --- a/tests/mcp_tests/test_aresponses_api_with_mcp.py +++ b/tests/mcp_tests/test_aresponses_api_with_mcp.py @@ -86,6 +86,15 @@ async def test_mcp_helper_methods(): LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_always) == False ) + # A single approval-required reference must disable auto-execution for the + # whole request; otherwise a "never" reference alongside an "always" one + # would let the approval-gated tool run without approval. + mcp_tools_mixed = [{"require_approval": "never"}, {"require_approval": "always"}] + assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_mixed) == False + mcp_tools_manual = [{"require_approval": "never"}, {"require_approval": "manual"}] + assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools(mcp_tools_manual) == False + assert LiteLLM_Proxy_MCP_Handler._should_auto_execute_tools([]) == False + print("✓ MCP helper methods test passed!") diff --git a/tests/test_litellm/caching/test_disk_cache.py b/tests/test_litellm/caching/test_disk_cache.py index b8d3b7b8d36..084370726b1 100644 --- a/tests/test_litellm/caching/test_disk_cache.py +++ b/tests/test_litellm/caching/test_disk_cache.py @@ -1,3 +1,7 @@ +import threading +import time +from concurrent.futures import ThreadPoolExecutor + import pytest pytest.importorskip("diskcache") @@ -5,6 +9,12 @@ pytest.importorskip("diskcache") from litellm.caching.disk_cache import DiskCache +class _SlowInt(int): + def __add__(self, value: int) -> "_SlowInt": + time.sleep(0.05) + return _SlowInt(int(self) + value) + + @pytest.fixture def cache(tmp_path): return DiskCache(disk_cache_dir=str(tmp_path)) @@ -27,6 +37,22 @@ def test_increment_cache_treats_non_int_cached_value_as_zero(cache): assert cache.get_cache("counter") == 4 +def test_increment_cache_is_atomic_under_thread_concurrency(cache): + seed = 1000 + cache.set_cache("counter", _SlowInt(seed)) + thread_count = 8 + barrier = threading.Barrier(thread_count) + + def increment(_: int) -> int: + barrier.wait() + return cache.increment_cache("counter", 1) + + with ThreadPoolExecutor(max_workers=thread_count) as executor: + tuple(executor.map(increment, range(thread_count))) + + assert cache.get_cache("counter") == seed + thread_count + + async def test_async_increment_starts_from_zero_when_key_missing(cache): assert await cache.async_increment("counter", 2) == 2 diff --git a/tests/test_litellm/caching/test_in_memory_cache.py b/tests/test_litellm/caching/test_in_memory_cache.py index 8828ebf207e..7be03d23fbe 100644 --- a/tests/test_litellm/caching/test_in_memory_cache.py +++ b/tests/test_litellm/caching/test_in_memory_cache.py @@ -2,7 +2,9 @@ import asyncio import json import os import sys +import threading import time +from concurrent.futures import ThreadPoolExecutor from unittest.mock import MagicMock, patch import httpx @@ -18,6 +20,36 @@ from unittest.mock import AsyncMock from litellm.caching.in_memory_cache import InMemoryCache +class _SlowInt(int): + def __add__(self, value: int) -> "_SlowInt": + time.sleep(0.05) + return _SlowInt(int(self) + value) + + +def test_increment_cache_is_atomic_under_thread_concurrency(): + cache = InMemoryCache() + seed = 1000 + cache.set_cache("counter", _SlowInt(seed)) + thread_count = 8 + barrier = threading.Barrier(thread_count) + + def increment(_: int) -> float: + barrier.wait() + return cache.increment_cache("counter", 1) + + with ThreadPoolExecutor(max_workers=thread_count) as executor: + tuple(executor.map(increment, range(thread_count))) + + assert cache.get_cache("counter") == seed + thread_count + + +async def test_async_increment_delegates_to_locked_sync_path(): + cache = InMemoryCache() + assert await cache.async_increment("counter", 2) == 2 + assert await cache.async_increment("counter", 3) == 5 + assert cache.get_cache("counter") == 5 + + def test_in_memory_openai_obj_cache(): from openai import OpenAI diff --git a/tests/test_litellm/experimental_mcp_client/test_tools.py b/tests/test_litellm/experimental_mcp_client/test_tools.py index 786bbf7dcc9..804e99b6f4e 100644 --- a/tests/test_litellm/experimental_mcp_client/test_tools.py +++ b/tests/test_litellm/experimental_mcp_client/test_tools.py @@ -18,6 +18,7 @@ from mcp.types import ( from mcp.types import Tool as MCPTool from litellm.experimental_mcp_client.tools import ( + transform_mcp_tool_to_anthropic_tool, _get_function_arguments, _normalize_mcp_input_schema, call_mcp_tool, @@ -250,3 +251,93 @@ def test_transform_mcp_tool_to_openai_responses_api_tool(): assert "query" in openai_tool["parameters"]["properties"] assert openai_tool["parameters"]["required"] == ["query"] assert openai_tool["parameters"]["additionalProperties"] == False + + +def test_transform_mcp_tool_to_anthropic_tool(): + """ + Regression test (LIT-4517): MCP tools must reach /v1/messages in Anthropic's + own tool shape. + + Given: An MCP tool + When: It is transformed for the Anthropic Messages API + Then: It carries name/description/input_schema, the shape that endpoint + accepts, rather than an OpenAI function block + + /v1/messages rejects an OpenAI-shaped tool outright ("Input tag 'function' + does not match any of the expected tags"), so reusing either OpenAI + transform here loses every MCP tool. + """ + tool = MCPTool( + name="read_wiki_structure", + description="Get a list of documentation topics", + inputSchema={ + "type": "object", + "properties": {"repoName": {"type": "string"}}, + "required": ["repoName"], + }, + ) + + anthropic_tool = transform_mcp_tool_to_anthropic_tool(tool) + + assert anthropic_tool["name"] == "read_wiki_structure" + assert anthropic_tool["description"] == "Get a list of documentation topics" + assert anthropic_tool["type"] == "custom" + assert anthropic_tool["input_schema"]["type"] == "object" + assert "repoName" in anthropic_tool["input_schema"]["properties"] + assert anthropic_tool["input_schema"]["required"] == ["repoName"] + assert "function" not in anthropic_tool, "Anthropic tools must not carry an OpenAI function block" + assert "parameters" not in anthropic_tool, "Anthropic names the schema input_schema, not parameters" + + +def test_transform_mcp_tool_to_anthropic_tool_normalizes_empty_schema(): + """A tool with no declared arguments must still present a valid object schema.""" + anthropic_tool = transform_mcp_tool_to_anthropic_tool( + MCPTool(name="noargs", description=None, inputSchema={}) + ) + + assert anthropic_tool["name"] == "noargs" + assert anthropic_tool["description"] == "" + assert anthropic_tool["input_schema"]["type"] == "object" + assert anthropic_tool["input_schema"]["properties"] == {} + + +def test_transform_mcp_tool_to_anthropic_tool_strips_keys_anthropic_rejects(): + """ + Regression test (LIT-4517): an MCP schema with keys Anthropic does not accept + must be sanitized, so the same tool cannot succeed on /chat/completions and 400 + on /v1/messages. + + Given: An MCP tool whose inputSchema carries $schema, legacy definitions and oneOf + When: It is transformed for the Anthropic Messages API + Then: Only keys in AnthropicInputSchema survive, matching the chat path + + The chat path runs the schema through the same sanitizer, so before this the two + routes diverged: a clean-schema server (deepwiki) worked on both, but a server + with a richer schema would be rejected only on messages. + """ + from litellm.types.llms.anthropic import AnthropicInputSchema + + tool = MCPTool( + name="rich", + description="tool with a dirty schema", + inputSchema={ + "type": "object", + "properties": {"q": {"type": "string"}}, + "required": ["q"], + "$schema": "http://json-schema.org/draft-07/schema#", + "definitions": {"D": {"type": "string"}}, + "oneOf": [{"required": ["q"]}], + }, + ) + + anthropic_tool = transform_mcp_tool_to_anthropic_tool(tool) + schema_keys = set(anthropic_tool["input_schema"].keys()) + + assert schema_keys <= set(AnthropicInputSchema.__annotations__.keys()), ( + f"schema must only carry keys Anthropic accepts, got {schema_keys}" + ) + assert "$schema" not in schema_keys + assert "definitions" not in schema_keys + assert "oneOf" not in schema_keys + assert anthropic_tool["input_schema"]["properties"] == {"q": {"type": "string"}} + assert anthropic_tool["input_schema"]["required"] == ["q"] diff --git a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py index 70c1f65b541..d94f0d5f47e 100644 --- a/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py +++ b/tests/test_litellm/integrations/test_anthropic_cache_control_hook.py @@ -1265,8 +1265,11 @@ def test_cache_control_hook_reserves_slot_for_tool_config_point(): ) assert _count_cache_control(processed) == 3 - # The tool_config point is passed through for the provider transform. - assert non_default_params["cache_control_injection_points"] == [{"location": "tool_config"}] + # The tool_config point is passed through for the provider transform, + # stamped so re-entries never re-judge it against litellm's own marks. + assert non_default_params["cache_control_injection_points"] == [ + {"location": "tool_config", "_litellm_judged": True} + ] @pytest.mark.asyncio @@ -1622,6 +1625,13 @@ class TestEnableAnthropicPromptCaching: monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) assert [p["index"] for p in self._points(tools=tools)] == [None, -1] + def test_stands_down_when_tool_function_carries_cache_control(self, monkeypatch): + """OpenAI-shaped tools nest cache_control under ``function``; the Anthropic + chat transform honors that location, so the stand-down must see it too.""" + monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) + tools = [{"type": "function", "function": {"name": "t", "parameters": {}, "cache_control": {"type": "ephemeral"}}}] + assert self._points(tools=tools) == [] + def test_seed_stands_down_when_only_tools_carry_cache_control(self, monkeypatch): """Same guard on the /chat/completions seeding path.""" monkeypatch.setattr(litellm, "enable_anthropic_prompt_caching", True) @@ -1726,6 +1736,131 @@ class TestEnableAnthropicPromptCaching: assert result_msgs == messages +class TestConfiguredInjectionPointsStandDown: + """Configured cache_control_injection_points must stand down entirely when the + client already set its own cache_control anywhere in the request (LIT-4582); + injecting alongside client breakpoints clashes with the client's caching + strategy and can push the request past Anthropic's four-block limit.""" + + CONFIGURED = [{"location": "message", "role": "system"}] + + CLEAN_MESSAGES: List[AllMessageValues] = [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": "hi"}, + ] + + MARKED_MESSAGES: List[AllMessageValues] = [ + {"role": "system", "content": "sys"}, + {"role": "user", "content": [{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}}]}, + ] + + V1_MESSAGES = [{"role": "user", "content": [{"type": "text", "text": "hi"}]}] + + def _seed(self, params, messages, tools=None): + AnthropicCacheControlHook.maybe_seed_default_injection_points( + non_default_params=params, + messages=messages, + model="claude-sonnet-4-5", + custom_llm_provider="anthropic", + tools=tools, + ) + + def _inject(self, messages, kwargs, system="sys", tools=None): + return AnthropicCacheControlHook.maybe_inject_cache_control( + messages, + system, + kwargs, + model="claude-sonnet-4-5", + custom_llm_provider="anthropic", + tools=tools, + ) + + def test_configured_points_dropped_when_messages_carry_cache_control(self): + params = {"cache_control_injection_points": copy.deepcopy(self.CONFIGURED)} + self._seed(params, copy.deepcopy(self.MARKED_MESSAGES)) + assert "cache_control_injection_points" not in params + + @pytest.mark.parametrize( + "tool", + [ + {"type": "function", "function": {"name": "t", "parameters": {}}, "cache_control": {"type": "ephemeral"}}, + {"type": "function", "function": {"name": "t", "parameters": {}, "cache_control": {"type": "ephemeral"}}}, + ], + ids=["top_level", "nested_in_function"], + ) + def test_configured_points_dropped_when_tools_carry_cache_control(self, tool): + params = {"cache_control_injection_points": copy.deepcopy(self.CONFIGURED)} + self._seed(params, copy.deepcopy(self.CLEAN_MESSAGES), tools=[tool]) + assert "cache_control_injection_points" not in params + + def test_configured_points_kept_when_request_is_unmarked(self): + configured = copy.deepcopy(self.CONFIGURED) + params = {"cache_control_injection_points": configured} + self._seed(params, copy.deepcopy(self.CLEAN_MESSAGES)) + assert params["cache_control_injection_points"] is configured + + def test_judged_remainder_survives_reentry_despite_injected_marks(self): + """acompletion() re-enters completion() after injection ran, with only the + stamped non-message points written back; the re-entry must not misread + litellm's own marks as client ones and drop that remainder.""" + remainder = [{"location": "tool_config", "_litellm_judged": True}] + params = {"cache_control_injection_points": remainder} + self._seed(params, copy.deepcopy(self.MARKED_MESSAGES)) + assert params["cache_control_injection_points"] is remainder + + def test_v1_messages_stand_down_when_content_block_marked(self): + messages = [ + {"role": "user", "content": [{"type": "text", "text": "hi", "cache_control": {"type": "ephemeral"}}]} + ] + kwargs = {"cache_control_injection_points": copy.deepcopy(self.CONFIGURED)} + result_msgs, result_sys = self._inject(copy.deepcopy(messages), kwargs) + assert result_msgs == messages + assert result_sys == "sys" + assert "cache_control_injection_points" not in kwargs + + def test_v1_messages_stand_down_when_system_block_marked(self): + """A configured point targeting a message must not fire when the client + marked the system prompt; the old behavior injected into the message + because only the exact targeted position was guarded.""" + system = [{"type": "text", "text": "s", "cache_control": {"type": "ephemeral"}}] + kwargs = {"cache_control_injection_points": [{"location": "message", "role": "user"}]} + result_msgs, result_sys = self._inject(copy.deepcopy(self.V1_MESSAGES), kwargs, system=system) + assert result_msgs == self.V1_MESSAGES + assert result_sys == system + assert "cache_control_injection_points" not in kwargs + + def test_v1_messages_stand_down_when_tools_marked(self): + tools = [{"name": "t", "input_schema": {}, "cache_control": {"type": "ephemeral"}}] + kwargs = {"cache_control_injection_points": copy.deepcopy(self.CONFIGURED)} + result_msgs, result_sys = self._inject(copy.deepcopy(self.V1_MESSAGES), kwargs, tools=tools) + assert result_msgs == self.V1_MESSAGES + assert result_sys == "sys" + assert "cache_control_injection_points" not in kwargs + + def test_v1_messages_configured_points_apply_when_unmarked(self): + kwargs = {"cache_control_injection_points": copy.deepcopy(self.CONFIGURED)} + _, result_sys = self._inject(copy.deepcopy(self.V1_MESSAGES), kwargs) + assert result_sys == [{"type": "text", "text": "sys", "cache_control": {"type": "ephemeral"}}] + + def test_v1_messages_reentry_flow_preserves_tool_config_remainder(self): + """The advisor interceptor re-enters anthropic_messages() with the outer + request's kwargs and post-injection messages. The first pass applies the + message point and writes back a stamped tool_config remainder; the + re-entry must keep that remainder even though the messages and system + now carry litellm's own marks.""" + points = [{"location": "message", "role": "system"}, {"location": "tool_config"}] + kwargs = {"cache_control_injection_points": copy.deepcopy(points)} + msgs1, sys1 = self._inject(copy.deepcopy(self.V1_MESSAGES), kwargs) + assert sys1[0]["cache_control"] == {"type": "ephemeral"} + expected_remainder = [{"location": "tool_config", "_litellm_judged": True}] + assert kwargs["cache_control_injection_points"] == expected_remainder + + msgs2, sys2 = self._inject(msgs1, kwargs, system=sys1) + assert kwargs["cache_control_injection_points"] == expected_remainder + assert msgs2 == msgs1 + assert sys2 == sys1 + + class TestAnthropicPromptCachingEnvVars: """Both settings are read from the environment at import, so an admin can enable auto-caching without a config file. Each case re-imports litellm in a subprocess diff --git a/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py new file mode 100644 index 00000000000..060c3e459d0 --- /dev/null +++ b/tests/test_litellm/llms/anthropic/experimental_pass_through/messages/test_mcp_handler.py @@ -0,0 +1,243 @@ +import os +import sys +from unittest.mock import AsyncMock, patch + +import pytest + +sys.path.insert(0, os.path.abspath("../../../../../..")) + +from litellm.llms.anthropic.experimental_pass_through.messages.handler import ( + anthropic_messages_handler, +) +from litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler import ( + _build_tool_result_message, + _extract_tool_use_blocks, +) + +MCP_REFERENCE = { + "type": "mcp", + "server_label": "litellm", + "server_url": "litellm_proxy/mcp/deepwiki", + "require_approval": "never", +} + + +def test_anthropic_messages_handler_routes_litellm_proxy_mcp_to_the_gateway(): + """ + Regression test (LIT-4517): /v1/messages must expand a litellm_proxy MCP + reference through the MCP gateway. + + Given: A /v1/messages request whose tools carry a litellm_proxy MCP reference + When: The handler dispatches + Then: It hands off to the MCP gateway instead of the provider + + Without this hook the reference is forwarded to Anthropic verbatim and the API + rejects the request ("Input tag 'mcp' found using 'type' does not match any of + the expected tags"), because only /v1/chat/completions and /v1/responses ever + had a gateway entry point. This pins the wiring, not the helper: deleting the + dispatch makes the whole feature unreachable while every unit test still passes. + """ + with patch( + "litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp", + new=AsyncMock(return_value={"routed": True}), + ) as routed: + result = anthropic_messages_handler( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="claude-sonnet-4-5", + tools=[MCP_REFERENCE], + custom_llm_provider="anthropic", + ) + + assert routed.called, "A litellm_proxy MCP reference must be dispatched to the MCP gateway" + assert routed.call_args.kwargs["tools"] == [MCP_REFERENCE] + assert routed.call_args.kwargs["model"] == "claude-sonnet-4-5" + assert result is not None + + +def test_anthropic_messages_handler_skips_the_gateway_on_recursion(): + """The gateway's own follow-up call must not re-enter the gateway.""" + with patch( + "litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp", + new=AsyncMock(return_value={"routed": True}), + ) as routed: + with pytest.raises(Exception): + anthropic_messages_handler( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="claude-sonnet-4-5", + tools=[MCP_REFERENCE], + custom_llm_provider="anthropic", + _skip_mcp_handler=True, + ) + + assert not routed.called, "_skip_mcp_handler must stop the gateway from recursing" + + +def test_anthropic_messages_handler_leaves_native_tools_alone(): + """A plain Anthropic tool is not an MCP reference and must not reach the gateway.""" + with patch( + "litellm.llms.anthropic.experimental_pass_through.messages.mcp_handler.anthropic_messages_with_mcp", + new=AsyncMock(return_value={"routed": True}), + ) as routed: + with pytest.raises(Exception): + anthropic_messages_handler( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="claude-sonnet-4-5", + tools=[{"name": "get_weather", "input_schema": {"type": "object"}}], + custom_llm_provider="anthropic", + ) + + assert not routed.called, "Only litellm_proxy MCP references belong to the gateway" + + +def test_extract_tool_use_blocks_ignores_text_blocks(): + """Only tool_use blocks drive the loop; text blocks are the model's prose.""" + response = { + "content": [ + {"type": "text", "text": "let me look that up"}, + {"type": "tool_use", "id": "toolu_1", "name": "read_wiki_structure", "input": {"repoName": "a/b"}}, + ] + } + + blocks = _extract_tool_use_blocks(response) + + assert len(blocks) == 1 + assert blocks[0]["name"] == "read_wiki_structure" + + +def test_build_tool_result_message_uses_anthropic_tool_result_blocks(): + """ + Results must go back as tool_result blocks in a user message. + + Anthropic pairs each result to its request by tool_use_id; the OpenAI shape + (a role="tool" message keyed by tool_call_id) is rejected here. + """ + message = _build_tool_result_message([{"tool_call_id": "toolu_1", "result": "9 sections", "name": "read_wiki"}]) + + assert message["role"] == "user" + assert list(message["content"]) == [ + {"type": "tool_result", "tool_use_id": "toolu_1", "content": "9 sections"} + ] + + +@pytest.mark.asyncio +async def test_anthropic_messages_with_mcp_forwards_the_callers_mcp_credentials(): + """ + Regression test (LIT-4517): the caller's MCP auth must reach both tool listing + and tool execution on /v1/messages. + + Given: A request carrying MCP auth headers and request tags + When: The gateway lists and then executes an MCP tool + Then: Both calls receive the caller's credentials, tags and trace ids + + Dropping them does not fail loudly; the tool still executes, just with no + credentials, so every auth-requiring MCP server (interactive OAuth, bearer + token, per-user env) silently returns nothing while the model claims it has + no access. Only a no-auth server would look healthy. + """ + from litellm.llms.anthropic.experimental_pass_through.messages import mcp_handler + from litellm.responses.mcp.request_context import MCPRequestContext + + context = MCPRequestContext( + user_api_key_auth="auth-object", + mcp_auth_header="legacy-header", + mcp_server_auth_headers={"deepwiki": {"authorization": "Bearer per-server"}}, + oauth2_headers={"authorization": "Bearer oauth"}, + raw_headers={"x-trace": "abc"}, + request_tags=["team-a"], + litellm_trace_id="trace-123", + litellm_call_id="call-456", + ) + + process = AsyncMock(return_value=([], {})) + execute = AsyncMock(return_value=[{"tool_call_id": "toolu_1", "result": "ok", "name": "t"}]) + responses = [ + {"stop_reason": "tool_use", "content": [{"type": "tool_use", "id": "toolu_1", "name": "t", "input": {}}]}, + {"stop_reason": "end_turn", "content": [{"type": "text", "text": "done"}]}, + ] + + with patch.object(MCPRequestContext, "resolve", return_value=context), patch.object( + mcp_handler.LiteLLM_Proxy_MCP_Handler + if hasattr(mcp_handler, "LiteLLM_Proxy_MCP_Handler") + else __import__( + "litellm.responses.mcp.litellm_proxy_mcp_handler", fromlist=["LiteLLM_Proxy_MCP_Handler"] + ).LiteLLM_Proxy_MCP_Handler, + "_process_mcp_tools_without_openai_transform", + new=process, + ), patch( + "litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._execute_tool_calls", + new=execute, + ), patch( + "litellm.anthropic_messages", new=AsyncMock(side_effect=responses) + ): + await mcp_handler.anthropic_messages_with_mcp( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="claude-sonnet-4-5", + tools=[MCP_REFERENCE], + ) + + listing = process.call_args.kwargs + assert listing["mcp_auth_header"] == "legacy-header", "tool listing must use the caller's MCP auth" + assert listing["mcp_server_auth_headers"] == {"deepwiki": {"authorization": "Bearer per-server"}} + assert listing["request_tags"] == ["team-a"] + assert listing["litellm_trace_id"] == "trace-123" + + execution = execute.call_args.kwargs + assert execution["user_api_key_auth"] == "auth-object" + assert execution["mcp_auth_header"] == "legacy-header", "tool execution must use the caller's MCP auth" + assert execution["mcp_server_auth_headers"] == {"deepwiki": {"authorization": "Bearer per-server"}} + assert execution["oauth2_headers"] == {"authorization": "Bearer oauth"} + assert execution["raw_headers"] == {"x-trace": "abc"} + assert execution["litellm_call_id"] == "call-456" + assert execution["litellm_trace_id"] == "trace-123" + assert execution["request_tags"] == ["team-a"] + + +@pytest.mark.asyncio +async def test_anthropic_messages_with_mcp_stops_when_every_tool_call_is_skipped(): + """ + Regression test (LIT-4517): a tool_use turn whose calls all get skipped must + end the loop, not send an empty tool_result message. + + Given: The model asks for a tool but the executor skips it (unresolvable name) + When: The gateway loop handles the empty result set + Then: It returns the last response instead of calling the model again + + _build_tool_result_message([]) produces a user message with empty content, and + Anthropic rejects that, so the caller would get an unhandled 400 from the middle + of the loop rather than the model's own answer. + """ + from litellm.llms.anthropic.experimental_pass_through.messages import mcp_handler + from litellm.responses.mcp.request_context import MCPRequestContext + + tool_use_response = { + "stop_reason": "tool_use", + "content": [{"type": "tool_use", "id": "toolu_1", "name": "gone", "input": {}}], + } + anthropic_messages_mock = AsyncMock(return_value=tool_use_response) + + with patch.object( + MCPRequestContext, "resolve", return_value=MCPRequestContext(user_api_key_auth="auth") + ), patch( + "litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._process_mcp_tools_without_openai_transform", + new=AsyncMock(return_value=([], {})), + ), patch( + "litellm.responses.mcp.litellm_proxy_mcp_handler.LiteLLM_Proxy_MCP_Handler._execute_tool_calls", + new=AsyncMock(return_value=[]), + ), patch( + "litellm.anthropic_messages", new=anthropic_messages_mock + ): + result = await mcp_handler.anthropic_messages_with_mcp( + max_tokens=100, + messages=[{"role": "user", "content": "hi"}], + model="claude-sonnet-4-5", + tools=[MCP_REFERENCE], + ) + + assert anthropic_messages_mock.await_count == 1, ( + "With no tool results there is nothing to send back, so the loop must not call the model again" + ) + assert result == tool_use_response diff --git a/tests/test_litellm/llms/bedrock/batches/test_transformation.py b/tests/test_litellm/llms/bedrock/batches/test_transformation.py index d1ad5943ae6..3681daffe5e 100644 --- a/tests/test_litellm/llms/bedrock/batches/test_transformation.py +++ b/tests/test_litellm/llms/bedrock/batches/test_transformation.py @@ -258,6 +258,112 @@ def test_create_request_no_timeout_for_non_24h_window(config): assert "timeoutDurationInHours" not in mock_sign.call_args.kwargs["data"] +def test_create_request_forwards_bedrock_tags_from_litellm_params(config): + tags = [ + {"key": "application", "value": "genai-proxy"}, + {"key": "team", "value": "ml-platform"}, + ] + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={ + "aws_batch_role_arn": "arn:aws:iam::1:role/r", + "bedrock_tags": tags, + }, + ) + assert mock_sign.call_args.kwargs["data"]["tags"] == tags + + +def test_create_request_forwards_bedrock_tags_from_optional_params(config): + tags = [{"key": "env", "value": "prod"}] + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={"bedrock_tags": tags}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + assert mock_sign.call_args.kwargs["data"]["tags"] == tags + + +def test_create_request_empty_litellm_params_tags_do_not_fall_through(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={"bedrock_tags": [{"key": "env", "value": "prod"}]}, + litellm_params={ + "aws_batch_role_arn": "arn:aws:iam::1:role/r", + "bedrock_tags": [], + }, + ) + assert mock_sign.call_args.kwargs["data"]["tags"] == [] + + +def test_create_request_omits_tags_when_bedrock_tags_absent(config): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={"aws_batch_role_arn": "arn:aws:iam::1:role/r"}, + ) + assert "tags" not in mock_sign.call_args.kwargs["data"] + + +@pytest.mark.parametrize( + "bad_tags", + [ + ["application=genai-proxy"], + [{"key": "application"}], + [{"value": "genai-proxy"}], + [{"key": "application", "value": 42}], + {"key": "application", "value": "genai-proxy"}, + "application=genai-proxy", + ], +) +def test_create_request_rejects_malformed_bedrock_tags(config, bad_tags): + with patch.object( + config.common_utils, + "generate_unique_job_name", + return_value="litellm-batch-1", + ), patch.object(config.common_utils, "sign_aws_request") as mock_sign: + mock_sign.return_value = ({}, b"{}") + with pytest.raises(ValueError, match="Invalid 'bedrock_tags' value"): + config.transform_create_batch_request( + model="m", + create_batch_data={"input_file_id": "s3://b/in.jsonl"}, + optional_params={}, + litellm_params={ + "aws_batch_role_arn": "arn:aws:iam::1:role/r", + "bedrock_tags": bad_tags, + }, + ) + mock_sign.assert_not_called() + + # --------------------------------------------------------------------------- # # transform_create_batch_response - status mapping + LiteLLMBatch shape # --------------------------------------------------------------------------- # diff --git a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py index 6809799d34f..94945ed4bfb 100644 --- a/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py +++ b/tests/test_litellm/llms/fireworks_ai/chat/test_fireworks_ai_chat_transformation.py @@ -123,6 +123,90 @@ def test_validate_environment_preserves_explicit_session_affinity_header(): assert headers["x-session-affinity"] == "explicit-session" +def test_validate_environment_sets_json_content_type(): + config = FireworksAIConfig() + + headers = config.validate_environment( + headers={}, + model="accounts/fireworks/models/test-model", + messages=[], + optional_params={}, + litellm_params={}, + api_key="test-key", + ) + + assert headers["Content-Type"] == "application/json" + + +def test_validate_environment_preserves_explicit_content_type(): + config = FireworksAIConfig() + + headers = config.validate_environment( + headers={"content-type": "multipart/form-data"}, + model="accounts/fireworks/models/test-model", + messages=[], + optional_params={}, + litellm_params={}, + api_key="test-key", + ) + + assert headers["content-type"] == "multipart/form-data" + assert "Content-Type" not in headers + + +def test_validate_environment_sets_json_content_type_with_session_affinity(): + config = FireworksAIConfig() + + headers = config.validate_environment( + headers={}, + model="accounts/fireworks/models/test-model", + messages=[], + optional_params={}, + litellm_params={"litellm_session_id": "session-123"}, + api_key="test-key", + ) + + assert headers["Content-Type"] == "application/json" + assert headers["Authorization"] == "Bearer test-key" + assert headers["x-session-affinity"] == "session-123" + + +def test_validate_environment_resolves_api_key_from_env_and_sets_content_type(monkeypatch): + monkeypatch.setenv("FIREWORKS_API_KEY", "fw-env-key") + config = FireworksAIConfig() + + headers = config.validate_environment( + headers={}, + model="accounts/fireworks/models/test-model", + messages=[], + optional_params={}, + litellm_params={}, + ) + + assert headers["Authorization"] == "Bearer fw-env-key" + assert headers["Content-Type"] == "application/json" + + +def test_validate_environment_raises_without_api_key(monkeypatch): + for env_var in ( + "FIREWORKS_API_KEY", + "FIREWORKS_AI_API_KEY", + "FIREWORKSAI_API_KEY", + "FIREWORKS_AI_TOKEN", + ): + monkeypatch.delenv(env_var, raising=False) + config = FireworksAIConfig() + + with pytest.raises(ValueError, match="FIREWORKS_API_KEY is not set"): + config.validate_environment( + headers={}, + model="accounts/fireworks/models/test-model", + messages=[], + optional_params={}, + litellm_params={}, + ) + + def test_get_fireworks_session_id_prefers_litellm_session_id_over_trace_id(): assert ( get_fireworks_session_id( diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py index 707374e7061..bf757b64c9a 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_adapter.py @@ -21,6 +21,7 @@ from litellm.proxy._experimental.mcp_server.outbound_credentials.adapter import from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( ApiKeyConfig, AuthorizationCodeConfig, + ClientCredentialsConfig, ClientSecretAuth, CredError, IdJagConfig, @@ -107,7 +108,6 @@ def test_oauth2_user_token_maps_to_authorization_code(oauth2_flow): [ _server(auth_type=MCPAuth.api_key), # no token configured _server(auth_type=MCPAuth.bearer_token), # no token configured - _server(auth_type=MCPAuth.oauth2, oauth2_flow="client_credentials"), # M2M -> v1 _server(auth_type=MCPAuth.oauth2, delegate_auth_to_upstream=True), # delegated upstream OAuth -> v1 _server(auth_type=MCPAuth.oauth2_token_exchange), # no endpoint/client creds -> incomplete -> v1 _server( @@ -124,6 +124,74 @@ def test_unmigrated_modes_defer_to_v1(server): assert to_server_spec(server) is None +def test_client_credentials_maps_full_config(): + spec = to_server_spec( + _server( + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + url="https://up.example.com/mcp", + token_url="https://idp.example.com/token", + client_id="cid", + client_secret="csec", + scopes=["read", "write"], + audience="https://up.example.com", + token_endpoint_auth_method="client_secret_basic", + ) + ) + assert spec is not None + config = spec.config + assert isinstance(config, ClientCredentialsConfig) + assert config.client_id == "cid" + assert config.client_secret is not None + assert config.client_secret.get_secret_value() == "csec" + assert config.token_url == "https://idp.example.com/token" + assert config.scopes == ("read", "write") + assert config.audience == "https://up.example.com" + assert config.token_endpoint_auth_method == "client_secret_basic" + + +def test_client_credentials_omits_audience_when_unset(): + spec = to_server_spec( + _server( + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + token_url="https://idp.example.com/token", + client_id="cid", + client_secret="csec", + ) + ) + assert spec is not None + assert isinstance(spec.config, ClientCredentialsConfig) + assert spec.config.audience is None + + +def test_client_credentials_with_incomplete_grant_fields_is_owned_for_fail_closed(): + # An M2M server missing its grant fields is still owned by v2 (spec, not None) so it fails + # closed at the source (misconfigured, 500) rather than deferring to v1, which would connect + # unauthenticated and mask the upstream 401 as an empty tool list. + spec = to_server_spec(_server(auth_type=MCPAuth.oauth2, oauth2_flow="client_credentials", client_id="cid")) + assert spec is not None + assert isinstance(spec.config, ClientCredentialsConfig) + assert spec.config.token_url is None + assert spec.config.client_secret is None + + +def test_client_credentials_wins_over_delegate_flag(): + # v1 never delegates for M2M servers; the explicit oauth2_flow opt-in outranks the delegate flag. + spec = to_server_spec( + _server( + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + delegate_auth_to_upstream=True, + token_url="https://idp.example.com/token", + client_id="cid", + client_secret="csec", + ) + ) + assert spec is not None + assert isinstance(spec.config, ClientCredentialsConfig) + + def test_token_exchange_maps_full_config(): spec = to_server_spec( _server( 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 new file mode 100644 index 00000000000..4e162090fbe --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_client_credentials.py @@ -0,0 +1,402 @@ +"""Tests for the client_credentials (M2M) token source and its retrying bearer auth. + +These are the behavior-contract spec: grant shape (scopes / audience / client auth method), +rotation-aware cache keying, expires_in-driven expiry, error classification, and the +401 -> discard -> refetch -> retry-once recovery in ``ClientCredentialsBearerAuth``. +""" + +import httpx +import pytest +from pydantic import SecretStr + +from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import ( + ClientCredentialsBearerAuth, + ClientCredentialsTokenSource, + TokenEndpointDenied, + TokenEndpointOutcome, + TokenEndpointSuccess, + TokenEndpointUnreachable, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Error, Ok +from litellm.proxy._experimental.mcp_server.outbound_credentials.types import ( + ClientCredentialsConfig, +) + + +class _Clock: + def __init__(self, t: float = 1000.0) -> None: + self.t = t + + def __call__(self) -> float: + return self.t + + +class _FakePoster: + """Records every grant POST and returns canned outcomes (last one repeats).""" + + def __init__(self, outcomes: "list[TokenEndpointOutcome]") -> None: + self._outcomes = outcomes + self.calls: "list[tuple[str, dict[str, str], dict[str, str]]]" = [] + + async def __call__(self, url: str, form: "dict[str, str]", headers: "dict[str, str]") -> TokenEndpointOutcome: + self.calls.append((url, dict(form), dict(headers))) + index = min(len(self.calls) - 1, len(self._outcomes) - 1) + return self._outcomes[index] + + +def _success(access_token: str = "m2m-token", **extra: object) -> TokenEndpointSuccess: + return TokenEndpointSuccess(body={"access_token": access_token, **extra}) + + +def _config(**overrides: object) -> ClientCredentialsConfig: + fields: "dict[str, object]" = { + "client_id": "cid", + "client_secret": SecretStr("csec"), + "token_url": "https://idp.example.com/token", + **overrides, + } + return ClientCredentialsConfig.model_validate(fields) + + +@pytest.mark.asyncio +async def test_grant_posts_client_credentials_with_scopes_and_audience(): + poster = _FakePoster([_success()]) + source = ClientCredentialsTokenSource(poster) + result = await source.get("s", _config(scopes=("read", "write"), audience="https://api.example.com")) + assert isinstance(result, Ok) + assert result.ok.access_token == "m2m-token" + url, form, _headers = poster.calls[0] + assert url == "https://idp.example.com/token" + assert form["grant_type"] == "client_credentials" + assert form["scope"] == "read write" + assert form["audience"] == "https://api.example.com" + assert form["client_id"] == "cid" + assert form["client_secret"] == "csec" + + +@pytest.mark.asyncio +async def test_grant_omits_scope_and_audience_when_not_configured(): + poster = _FakePoster([_success()]) + await ClientCredentialsTokenSource(poster).get("s", _config()) + _url, form, _headers = poster.calls[0] + assert "scope" not in form + assert "audience" not in form + + +@pytest.mark.asyncio +async def test_grant_honors_client_secret_basic(): + poster = _FakePoster([_success()]) + await ClientCredentialsTokenSource(poster).get("s", _config(token_endpoint_auth_method="client_secret_basic")) + _url, form, headers = poster.calls[0] + assert headers["Authorization"].startswith("Basic ") + assert "client_secret" not in form + assert "client_id" not in form + + +@pytest.mark.asyncio +async def test_missing_grant_fields_are_misconfigured_and_never_posted(): + poster = _FakePoster([_success()]) + result = await ClientCredentialsTokenSource(poster).get( + "s", ClientCredentialsConfig(client_id="cid", client_secret=SecretStr("csec")) + ) + assert isinstance(result, Error) + assert result.error.tag == "misconfigured" + assert "token_url" in result.error.summary + assert poster.calls == [] + + +@pytest.mark.asyncio +async def test_token_is_cached_across_gets(): + poster = _FakePoster([_success(expires_in=3600)]) + source = ClientCredentialsTokenSource(poster) + first = await source.get("s", _config()) + second = await source.get("s", _config()) + assert isinstance(first, Ok) and isinstance(second, Ok) + assert second.ok.access_token == first.ok.access_token + assert len(poster.calls) == 1 + + +@pytest.mark.asyncio +async def test_expires_in_bounds_the_cache_lifetime(): + clock = _Clock(1000.0) + poster = _FakePoster([_success("t1", expires_in=120), _success("t2", expires_in=120)]) + source = ClientCredentialsTokenSource(poster, expiry_skew_seconds=60.0, clock=clock) + first = await source.get("s", _config()) + assert isinstance(first, Ok) + assert first.ok.expires_at == 1120.0 + clock.t = 1059.0 # within expires_in - skew + assert len(poster.calls) == 1 + within = await source.get("s", _config()) + assert isinstance(within, Ok) and within.ok.access_token == "t1" + clock.t = 1061.0 # past expires_in - skew: the entry lapsed before the real token does + lapsed = await source.get("s", _config()) + assert isinstance(lapsed, Ok) and lapsed.ok.access_token == "t2" + assert len(poster.calls) == 2 + + +@pytest.mark.asyncio +async def test_short_lived_token_is_never_served_past_its_expiry(): + # expires_in below the skew must not be floored into serving an expired token: the cache + # entry lapses with the token itself, and the next get re-fetches. + clock = _Clock(1000.0) + poster = _FakePoster([_success("t1", expires_in=5), _success("t2", expires_in=5)]) + source = ClientCredentialsTokenSource(poster, expiry_skew_seconds=60.0, min_cache_seconds=10.0, clock=clock) + first = await source.get("s", _config()) + assert isinstance(first, Ok) and first.ok.access_token == "t1" + clock.t = 1004.0 # still within the token's real lifetime + within = await source.get("s", _config()) + assert isinstance(within, Ok) and within.ok.access_token == "t1" + clock.t = 1006.0 # past expires_at: the floor must not keep serving t1 + lapsed = await source.get("s", _config()) + assert isinstance(lapsed, Ok) and lapsed.ok.access_token == "t2" + assert len(poster.calls) == 2 + + +class _RecordingBackend: + """A TokenCacheBackend spy: records every write so a test can assert none happened.""" + + def __init__(self) -> None: + self.set_ttls: list[float] = [] + + async def get(self, identity_key: str, server_id: str): + return None + + async def set(self, identity_key: str, server_id: str, token, ttl_seconds: float) -> None: + self.set_ttls.append(ttl_seconds) + + async def delete(self, identity_key: str, server_id: str) -> None: + return None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("expires_in", [0, -30]) +async def test_non_positive_expires_in_writes_no_cache_entry(expires_in): + # A dead-on-arrival entry (ttl 0) must not be written at all: it can never be served, but it + # would occupy a slot in the bounded backend and could evict a live token. The mint itself + # still succeeds for the current request, and the next get re-fetches. + backend = _RecordingBackend() + poster = _FakePoster([_success("t1", expires_in=expires_in), _success("t2", expires_in=expires_in)]) + source = ClientCredentialsTokenSource(poster, backend=backend) + first = await source.get("s", _config()) + assert isinstance(first, Ok) and first.ok.access_token == "t1" + again = await source.get("s", _config()) + assert isinstance(again, Ok) and again.ok.access_token == "t2" + assert backend.set_ttls == [] + assert len(poster.calls) == 2 + + +@pytest.mark.asyncio +async def test_lock_dict_is_bounded_for_ephemeral_server_ids(): + poster = _FakePoster([_success()]) + source = ClientCredentialsTokenSource(poster, max_locks=8) + for index in range(20): + result = await source.get(f"ephemeral-{index}", _config()) + assert isinstance(result, Ok) + assert len(source._locks) <= 8 + + +@pytest.mark.asyncio +async def test_missing_expires_in_is_cached_briefly_not_an_hour(): + clock = _Clock(1000.0) + poster = _FakePoster([_success("t1"), _success("t2")]) + source = ClientCredentialsTokenSource(poster, default_ttl_seconds=300.0, clock=clock) + first = await source.get("s", _config()) + assert isinstance(first, Ok) + assert first.ok.expires_at is None + clock.t = 1301.0 # past the default TTL; v1 would still be serving its 3600s-cached token + second = await source.get("s", _config()) + assert isinstance(second, Ok) and second.ok.access_token == "t2" + assert len(poster.calls) == 2 + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + "rotation", + [ + {"client_secret": SecretStr("rotated")}, + {"client_id": "cid-2"}, + {"scopes": ("admin",)}, + {"audience": "https://other.example.com"}, + {"token_url": "https://idp2.example.com/token"}, + ], +) +async def test_credential_rotation_invalidates_the_cached_token(rotation): + poster = _FakePoster([_success("old", expires_in=3600), _success("new", expires_in=3600)]) + source = ClientCredentialsTokenSource(poster) + before = await source.get("s", _config()) + after = await source.get("s", _config(**rotation)) + assert isinstance(before, Ok) and before.ok.access_token == "old" + assert isinstance(after, Ok) and after.ok.access_token == "new" + assert len(poster.calls) == 2 + + +@pytest.mark.asyncio +async def test_idp_4xx_is_misconfigured_and_5xx_is_unavailable(): + denied = await ClientCredentialsTokenSource( + _FakePoster([TokenEndpointDenied(status_code=401, detail="HTTP 401")]) + ).get("s", _config()) + assert isinstance(denied, Error) and denied.error.tag == "misconfigured" + down = await ClientCredentialsTokenSource( + _FakePoster([TokenEndpointDenied(status_code=503, detail="HTTP 503")]) + ).get("s", _config()) + assert isinstance(down, Error) and down.error.tag == "upstream_unavailable" + unreachable = await ClientCredentialsTokenSource(_FakePoster([TokenEndpointUnreachable(detail="dns")])).get( + "s", _config() + ) + assert isinstance(unreachable, Error) and unreachable.error.tag == "upstream_unavailable" + + +@pytest.mark.asyncio +async def test_response_without_access_token_is_misconfigured(): + poster = _FakePoster([TokenEndpointSuccess(body={"token_type": "Bearer"})]) + result = await ClientCredentialsTokenSource(poster).get("s", _config()) + assert isinstance(result, Error) + assert result.error.tag == "misconfigured" + + +@pytest.mark.asyncio +async def test_error_results_are_not_cached(): + poster = _FakePoster([TokenEndpointUnreachable(detail="down"), _success("recovered")]) + source = ClientCredentialsTokenSource(poster) + first = await source.get("s", _config()) + second = await source.get("s", _config()) + assert isinstance(first, Error) + assert isinstance(second, Ok) and second.ok.access_token == "recovered" + + +@pytest.mark.asyncio +async def test_refetch_discards_the_failed_token_and_mints_a_fresh_one(): + poster = _FakePoster([_success("stale", expires_in=3600), _success("fresh", expires_in=3600)]) + source = ClientCredentialsTokenSource(poster) + first = await source.get("s", _config()) + assert isinstance(first, Ok) + fresh = await source.refetch("s", _config(), failed_access_token="stale") + assert fresh == "fresh" + assert len(poster.calls) == 2 + after = await source.get("s", _config()) + assert isinstance(after, Ok) and after.ok.access_token == "fresh" + assert len(poster.calls) == 2 + + +@pytest.mark.asyncio +async def test_refetch_reuses_a_concurrent_replacement_without_a_second_grant(): + poster = _FakePoster([_success("replacement", expires_in=3600)]) + source = ClientCredentialsTokenSource(poster) + seeded = await source.get("s", _config()) + assert isinstance(seeded, Ok) + result = await source.refetch("s", _config(), failed_access_token="some-older-token") + assert result == "replacement" + assert len(poster.calls) == 1 + + +@pytest.mark.asyncio +async def test_refetch_returns_none_when_the_grant_fails(): + poster = _FakePoster([_success("stale"), TokenEndpointUnreachable(detail="down")]) + source = ClientCredentialsTokenSource(poster) + await source.get("s", _config()) + assert await source.refetch("s", _config(), failed_access_token="stale") is None + + +def _upstream(responses: "list[httpx.Response]") -> "tuple[httpx.MockTransport, list[str]]": + # The auth flow re-yields the same Request object on retry, so snapshot the Authorization + # value per send; holding the Request would show the post-retry mutation for both entries. + seen: "list[str]" = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request.headers.get("Authorization", "")) + return responses[min(len(seen) - 1, len(responses) - 1)] + + return httpx.MockTransport(handler), seen + + +@pytest.mark.asyncio +async def test_bearer_auth_sends_the_token_and_leaves_a_success_alone(): + transport, seen = _upstream([httpx.Response(200)]) + + async def refetch(failed: str) -> "str | None": + raise AssertionError("must not refetch on success") + + auth = ClientCredentialsBearerAuth("m2m-token", refetch) + async with httpx.AsyncClient(transport=transport, auth=auth) as client: + response = await client.get("https://upstream.example.com/mcp") + assert response.status_code == 200 + assert seen == ["Bearer m2m-token"] + + +@pytest.mark.asyncio +async def test_bearer_auth_retries_a_401_once_with_a_fresh_token(): + transport, seen = _upstream([httpx.Response(401), httpx.Response(200)]) + refetched: "list[str]" = [] + + async def refetch(failed: str) -> "str | None": + refetched.append(failed) + return "fresh-token" + + auth = ClientCredentialsBearerAuth("stale-token", refetch) + async with httpx.AsyncClient(transport=transport, auth=auth) as client: + response = await client.get("https://upstream.example.com/mcp") + assert response.status_code == 200 + assert refetched == ["stale-token"] + assert seen == ["Bearer stale-token", "Bearer fresh-token"] + + +@pytest.mark.asyncio +async def test_bearer_auth_remembers_the_rotated_token_for_later_requests(): + # The auth object lives for the whole MCP session (it is the httpx client's auth), so after a + # 401 recovery it must send the fresh token first on subsequent requests; re-sending the + # rejected one would burn a 401 round trip and the single retry on every call. + transport, seen = _upstream([httpx.Response(401), httpx.Response(200), httpx.Response(200)]) + refetched: "list[str]" = [] + + async def refetch(failed: str) -> "str | None": + refetched.append(failed) + return "fresh-token" + + auth = ClientCredentialsBearerAuth("stale-token", refetch) + async with httpx.AsyncClient(transport=transport, auth=auth) as client: + first = await client.get("https://upstream.example.com/mcp") + second = await client.get("https://upstream.example.com/mcp") + assert first.status_code == 200 and second.status_code == 200 + assert refetched == ["stale-token"] + assert seen == ["Bearer stale-token", "Bearer fresh-token", "Bearer fresh-token"] + + +@pytest.mark.asyncio +async def test_bearer_auth_surfaces_the_401_when_the_refetch_fails(): + transport, seen = _upstream([httpx.Response(401)]) + + async def refetch(failed: str) -> "str | None": + return None + + auth = ClientCredentialsBearerAuth("stale-token", refetch) + async with httpx.AsyncClient(transport=transport, auth=auth) as client: + response = await client.get("https://upstream.example.com/mcp") + assert response.status_code == 401 + assert len(seen) == 1 + + +@pytest.mark.asyncio +async def test_bearer_auth_gives_up_after_a_second_401(): + transport, seen = _upstream([httpx.Response(401), httpx.Response(401)]) + refetched: "list[str]" = [] + + async def refetch(failed: str) -> "str | None": + refetched.append(failed) + return "fresh-token" + + auth = ClientCredentialsBearerAuth("stale-token", refetch) + async with httpx.AsyncClient(transport=transport, auth=auth) as client: + response = await client.get("https://upstream.example.com/mcp") + assert response.status_code == 401 + assert len(seen) == 2 + assert refetched == ["stale-token"] + + +def test_bearer_auth_rejects_sync_clients(): + async def refetch(failed: str) -> "str | None": + return None + + auth = ClientCredentialsBearerAuth("token", refetch) + with httpx.Client(transport=httpx.MockTransport(lambda request: httpx.Response(200)), auth=auth) as client: + with pytest.raises(RuntimeError): + client.get("https://upstream.example.com/mcp") diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py index ba7720ffd51..a710da81962 100644 --- a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_resolver.py @@ -1,9 +1,10 @@ """Tests for the resolver dispatch: live arms produce auth, stubbed arms fail closed. -`none`, `api_key` (shared-key source), `passthrough`, `authorization_code`, and `token_exchange` are -implemented; every other arm, plus the `api_key` BYOK source, returns a typed `not_implemented` error -until its mode lands. Parametrizing the stubs over one config each also guards reachability: a dropped -`case` would hit `assert_never` and raise instead of returning the stub. +`none`, `api_key` (shared-key source), `passthrough`, `authorization_code`, `token_exchange`, and +`client_credentials` are implemented; every other arm, plus the `api_key` BYOK source, returns a +typed `not_implemented` error until its mode lands. Parametrizing the stubs over one config each +also guards reachability: a dropped `case` would hit `assert_never` and raise instead of +returning the stub. """ import httpx @@ -324,9 +325,106 @@ async def test_passthrough_without_inbound_token_is_a_no_op(): assert isinstance(result.ok, NoOpAuth) +class _FakeM2MSource: + """A ClientCredentialsTokenSource returning a canned result and recording refetches.""" + + def __init__(self, result) -> None: + self._result = result + self.gets: list[str] = [] + self.refetches: list[tuple[str, str]] = [] + + async def get(self, server_id: str, config): + self.gets.append(server_id) + return self._result + + async def refetch(self, server_id: str, config, failed_access_token: str): + self.refetches.append((server_id, failed_access_token)) + return "fresh-m2m" + + +_M2M = ClientCredentialsConfig( + client_id="cid", + client_secret=SecretStr("csec"), + token_url="https://idp.example.com/token", +) + + +async def _emitted_async(auth: httpx.Auth, respond=None) -> tuple[httpx.Headers, list[httpx.Request]]: + """Drive the async auth flow one request at a time, replying via ``respond`` when given.""" + seen: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + seen.append(request) + return respond(request) if respond else httpx.Response(200) + + async with httpx.AsyncClient(transport=httpx.MockTransport(handler), auth=auth) as client: + await client.get("https://upstream.example.com/mcp") + return seen[-1].headers, seen + + +@pytest.mark.asyncio +async def test_client_credentials_emits_the_minted_bearer(): + source = _FakeM2MSource(Ok(OAuthToken(access_token="m2m-at"))) + result = await UpstreamCredentialProvider(client_credentials_source=source).resolve_credentials( + _SUBJECT, _spec(_M2M) + ) + assert isinstance(result, Ok) + headers, _ = await _emitted_async(result.ok) + assert headers["Authorization"] == "Bearer m2m-at" + assert source.gets == ["s"] + + +@pytest.mark.asyncio +async def test_client_credentials_ignores_the_subject(): + # The contract's no-user-context clause: every caller shares the one client identity. + source = _FakeM2MSource(Ok(OAuthToken(access_token="m2m-at"))) + provider = UpstreamCredentialProvider(client_credentials_source=source) + alice = await provider.resolve_credentials(Subject(tenant_id="t1", subject_id="alice"), _spec(_M2M)) + bob = await provider.resolve_credentials(Subject(tenant_id="t2", subject_id="bob"), _spec(_M2M)) + assert isinstance(alice, Ok) and isinstance(bob, Ok) + alice_headers, _ = await _emitted_async(alice.ok) + bob_headers, _ = await _emitted_async(bob.ok) + assert alice_headers["Authorization"] == bob_headers["Authorization"] == "Bearer m2m-at" + + +@pytest.mark.asyncio +async def test_client_credentials_auth_retries_a_401_through_the_source(): + source = _FakeM2MSource(Ok(OAuthToken(access_token="stale-at"))) + result = await UpstreamCredentialProvider(client_credentials_source=source).resolve_credentials( + _SUBJECT, _spec(_M2M) + ) + assert isinstance(result, Ok) + + def respond(request: httpx.Request) -> httpx.Response: + is_stale = request.headers["Authorization"] == "Bearer stale-at" + return httpx.Response(401) if is_stale else httpx.Response(200) + + headers, seen = await _emitted_async(result.ok, respond) + assert headers["Authorization"] == "Bearer fresh-m2m" + assert len(seen) == 2 + assert source.refetches == [("s", "stale-at")] + + +@pytest.mark.asyncio +async def test_client_credentials_propagates_the_source_error(): + source = _FakeM2MSource(Error(CredError.of_upstream_unavailable("idp down"))) + result = await UpstreamCredentialProvider(client_credentials_source=source).resolve_credentials( + _SUBJECT, _spec(_M2M) + ) + assert isinstance(result, Error) + assert result.error.tag == "upstream_unavailable" + + +@pytest.mark.asyncio +async def test_client_credentials_with_no_source_wired_fails_closed_on_missing_config(): + # The default source validates the grant fields before any network is touched. + result = await UpstreamCredentialProvider().resolve_credentials(_SUBJECT, _spec(ClientCredentialsConfig())) + assert isinstance(result, Error) + assert result.error.tag == "misconfigured" + + _STUBBED = [ ("api_key_byok", ApiKeyConfig(key_source=Byok())), - ("client_credentials", ClientCredentialsConfig()), ("aws_sigv4", AwsSigV4Config(region="us-east-1")), ] diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_credentials.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_credentials.py new file mode 100644 index 00000000000..8fa7c15d2d3 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_credentials.py @@ -0,0 +1,135 @@ +"""Tests for the session-token KDF and the edge/token-endpoint resolvers.""" + +from datetime import datetime, timedelta, timezone + +import pytest + +from litellm.proxy._experimental.mcp_server.outbound_credentials.bridge_credentials import ( + envelope_keys_from_master_key, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_credentials import ( + NotSessionBearer, + SessionBearerAdmitted, + SessionBearerInvalid, + SessionRefreshInvalid, + SessionRefreshOpened, + is_session_bearer_shaped, + open_session_refresh_bearer, + resolve_session_bearer, + session_keys_from_master_key, +) +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + SESSION_TTL_SECONDS, + MintedSessionToken, + SessionPrincipal, + mint_session_refresh_token, + mint_session_token, +) + +NOW = datetime(2026, 1, 1, 12, 0, 0, tzinfo=timezone.utc) +MASTER_KEY = "sk-master-key-for-tests" +KEYS = session_keys_from_master_key(MASTER_KEY) +PRINCIPAL = SessionPrincipal(user_id="user-123", client_id="llm_client_abc") + + +def _access_token() -> str: + minted = mint_session_token(PRINCIPAL, KEYS, NOW) + assert isinstance(minted, MintedSessionToken) + return minted.token.get_secret_value() + + +def _refresh_token() -> str: + minted = mint_session_refresh_token(PRINCIPAL, KEYS, NOW) + assert isinstance(minted, MintedSessionToken) + return minted.token.get_secret_value() + + +def test_kdf_is_deterministic_and_key_length_is_256_bit(): + again = session_keys_from_master_key(MASTER_KEY) + assert again.signing_key.get_secret_value() == KEYS.signing_key.get_secret_value() + assert len(bytes.fromhex(KEYS.signing_key.get_secret_value())) == 32 + + +def test_kdf_domain_separated_from_envelope_keys(): + envelope_keys = envelope_keys_from_master_key(MASTER_KEY) + session_signing = KEYS.signing_key.get_secret_value() + assert session_signing != envelope_keys.signing_key.get_secret_value() + assert session_signing != envelope_keys.encryption_key.get_secret_value() + + +def test_kdf_differs_across_master_keys(): + other = session_keys_from_master_key("sk-a-different-master-key") + assert other.signing_key.get_secret_value() != KEYS.signing_key.get_secret_value() + + +@pytest.mark.parametrize( + "value,expected", + [ + ("Bearer sk-1234", False), + ("sk-1234", False), + ("Bearer llm_env_abc", False), + ("Bearer llm_refresh_abc", False), + ("llm_session_abc", True), + ("Bearer llm_session_abc", True), + ("bearer llm_srefresh_abc", True), + ], +) +def test_is_session_bearer_shaped(value, expected): + assert is_session_bearer_shaped(value) is expected + + +def test_resolve_admits_valid_access_token_with_and_without_scheme(): + token = _access_token() + for value in (token, f"Bearer {token}", f"bearer {token}"): + result = resolve_session_bearer(value, KEYS, NOW) + assert isinstance(result, SessionBearerAdmitted) + assert result.principal == PRINCIPAL + + +def test_resolve_passes_non_session_bearers_through(): + for value in ("Bearer sk-1234", "Bearer llm_env_whatever", "Bearer eyJhbGciOi"): + assert isinstance(resolve_session_bearer(value, KEYS, NOW), NotSessionBearer) + + +def test_resolve_fails_expired_token_closed_and_flags_expiry(): + token = _access_token() + later = NOW + timedelta(seconds=SESSION_TTL_SECONDS + 1) + result = resolve_session_bearer(f"Bearer {token}", KEYS, later) + assert isinstance(result, SessionBearerInvalid) + assert result.expired is True + + +def test_resolve_fails_tampered_token_closed_without_expiry_flag(): + token = _access_token() + tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb") + result = resolve_session_bearer(f"Bearer {tampered}", KEYS, NOW) + assert isinstance(result, SessionBearerInvalid) + assert result.expired is False + + +def test_resolve_rejects_refresh_token_at_the_edge(): + result = resolve_session_bearer(f"Bearer {_refresh_token()}", KEYS, NOW) + assert isinstance(result, SessionBearerInvalid) + assert result.expired is False + + +def test_resolve_wrong_master_key_fails_closed(): + other_keys = session_keys_from_master_key("sk-rotated-master-key") + result = resolve_session_bearer(f"Bearer {_access_token()}", other_keys, NOW) + assert isinstance(result, SessionBearerInvalid) + + +def test_refresh_grant_opens_for_the_issued_client(): + result = open_session_refresh_bearer(_refresh_token(), KEYS, NOW, expected_client_id="llm_client_abc") + assert isinstance(result, SessionRefreshOpened) + assert result.principal == PRINCIPAL + + +def test_refresh_grant_rejects_a_different_client(): + result = open_session_refresh_bearer(_refresh_token(), KEYS, NOW, expected_client_id="llm_client_other") + assert isinstance(result, SessionRefreshInvalid) + + +def test_refresh_grant_rejects_access_token_presented_as_refresh(): + result = open_session_refresh_bearer(_access_token(), KEYS, NOW, expected_client_id="llm_client_abc") + assert isinstance(result, SessionRefreshInvalid) diff --git a/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py new file mode 100644 index 00000000000..a43592ebe18 --- /dev/null +++ b/tests/test_litellm/proxy/_experimental/mcp_server/outbound_credentials/test_session_token.py @@ -0,0 +1,214 @@ +"""Tests for the identity-only gateway session token (mint/open, hostile-input totality).""" + +from datetime import datetime, timedelta, timezone + +import jwt +import pytest +from pydantic import SecretStr, ValidationError + +from litellm.proxy._experimental.mcp_server.outbound_credentials.session_token import ( + MAX_SESSION_TOKEN_BYTES, + SESSION_ISSUER, + SESSION_REFRESH_PREFIX, + SESSION_REFRESH_TTL_SECONDS, + SESSION_TOKEN_PREFIX, + SESSION_TTL_SECONDS, + MintedSessionToken, + NotASessionToken, + OpenedSessionToken, + SessionBadSignature, + SessionExpired, + SessionKeys, + SessionMalformed, + SessionPrincipal, + SessionTokenTooLarge, + is_session_refresh_token, + is_session_token, + mint_session_refresh_token, + mint_session_token, + open_session_refresh_token, + open_session_token, +) + +NOW = datetime(2026, 1, 1, 12, 0, 0, tzinfo=timezone.utc) +KEYS = SessionKeys(signing_key=SecretStr("k" * 32)) +OTHER_KEYS = SessionKeys(signing_key=SecretStr("x" * 32)) +PRINCIPAL = SessionPrincipal(user_id="user-123", client_id="llm_client_abc") + + +def _mint_access() -> str: + minted = mint_session_token(PRINCIPAL, KEYS, NOW) + assert isinstance(minted, MintedSessionToken) + return minted.token.get_secret_value() + + +def _mint_refresh() -> str: + minted = mint_session_refresh_token(PRINCIPAL, KEYS, NOW) + assert isinstance(minted, MintedSessionToken) + return minted.token.get_secret_value() + + +def _sign_claims(payload: dict, prefix: str = SESSION_TOKEN_PREFIX, keys: SessionKeys = KEYS) -> str: + return prefix + jwt.encode(payload, keys.signing_key.get_secret_value(), algorithm="HS256") + + +def _valid_claims(**overrides) -> dict: + base = { + "iss": SESSION_ISSUER, + "iat": int(NOW.timestamp()), + "exp": int((NOW + timedelta(seconds=600)).timestamp()), + "jti": "jti-fixed", + "kind": "session", + "user_id": "user-123", + "client_id": "llm_client_abc", + } + return {**base, **overrides} + + +def test_access_round_trip_recovers_principal_and_caps_ttl(): + minted = mint_session_token(PRINCIPAL, KEYS, NOW) + assert isinstance(minted, MintedSessionToken) + assert minted.expires_at == NOW + timedelta(seconds=SESSION_TTL_SECONDS) + token = minted.token.get_secret_value() + assert is_session_token(token) + assert not is_session_refresh_token(token) + opened = open_session_token(token, KEYS, NOW) + assert isinstance(opened, OpenedSessionToken) + assert opened.principal == PRINCIPAL + + +def test_refresh_round_trip_recovers_principal_and_caps_ttl(): + minted = mint_session_refresh_token(PRINCIPAL, KEYS, NOW) + assert isinstance(minted, MintedSessionToken) + assert minted.expires_at == NOW + timedelta(seconds=SESSION_REFRESH_TTL_SECONDS) + token = minted.token.get_secret_value() + assert is_session_refresh_token(token) + opened = open_session_refresh_token(token, KEYS, NOW) + assert isinstance(opened, OpenedSessionToken) + assert opened.principal == PRINCIPAL + + +def test_access_token_reprefixed_as_refresh_is_rejected_by_signed_kind(): + body = _mint_access().removeprefix(SESSION_TOKEN_PREFIX) + swapped = SESSION_REFRESH_PREFIX + body + assert isinstance(open_session_refresh_token(swapped, KEYS, NOW), SessionMalformed) + + +def test_refresh_token_reprefixed_as_access_is_rejected_by_signed_kind(): + body = _mint_refresh().removeprefix(SESSION_REFRESH_PREFIX) + swapped = SESSION_TOKEN_PREFIX + body + assert isinstance(open_session_token(swapped, KEYS, NOW), SessionMalformed) + + +def test_refresh_token_is_not_an_access_token_at_the_edge(): + assert isinstance(open_session_token(_mint_refresh(), KEYS, NOW), NotASessionToken) + + +def test_expired_access_token_is_expired_not_malformed(): + token = _mint_access() + at_expiry = NOW + timedelta(seconds=SESSION_TTL_SECONDS) + assert isinstance(open_session_token(token, KEYS, at_expiry), SessionExpired) + after = NOW + timedelta(seconds=SESSION_TTL_SECONDS + 1) + assert isinstance(open_session_token(token, KEYS, after), SessionExpired) + + +def test_still_valid_one_second_before_expiry(): + token = _mint_access() + just_before = NOW + timedelta(seconds=SESSION_TTL_SECONDS - 1) + assert isinstance(open_session_token(token, KEYS, just_before), OpenedSessionToken) + + +def test_tampered_signature_is_bad_signature(): + token = _mint_access() + tampered = token[:-2] + ("aa" if not token.endswith("aa") else "bb") + assert isinstance(open_session_token(tampered, KEYS, NOW), SessionBadSignature) + + +def test_key_rotation_invalidates_outstanding_tokens(): + token = _mint_access() + assert isinstance(open_session_token(token, OTHER_KEYS, NOW), SessionBadSignature) + + +@pytest.mark.parametrize( + "candidate,expected", + [ + ("sk-1234", NotASessionToken), + ("llm_env_something", NotASessionToken), + ("", NotASessionToken), + (SESSION_TOKEN_PREFIX, SessionMalformed), + (SESSION_TOKEN_PREFIX + "not-a-jwt", SessionMalformed), + (SESSION_TOKEN_PREFIX + "\ud800garbage", SessionMalformed), + (SESSION_TOKEN_PREFIX + "a" * (MAX_SESSION_TOKEN_BYTES + 1), SessionMalformed), + ], +) +def test_hostile_candidates_never_raise(candidate, expected): + assert isinstance(open_session_token(candidate, KEYS, NOW), expected) + + +def test_multibyte_candidate_over_byte_cap_but_under_char_cap_is_rejected(): + filler = "€" * (MAX_SESSION_TOKEN_BYTES // 3) + candidate = SESSION_TOKEN_PREFIX + filler + assert len(candidate) <= MAX_SESSION_TOKEN_BYTES + assert isinstance(open_session_token(candidate, KEYS, NOW), SessionMalformed) + + +def test_alg_none_token_is_rejected(): + unsigned = jwt.api_jws.encode(b'{"iss":"litellm-mcp-gateway"}', key=None, algorithm="none") + assert isinstance(open_session_token(SESSION_TOKEN_PREFIX + unsigned, KEYS, NOW), SessionMalformed) + + +@pytest.mark.parametrize( + "claims", + [ + _valid_claims(iss="wrong-issuer"), + _valid_claims(exp=str(int((NOW + timedelta(seconds=600)).timestamp()))), + _valid_claims(iat="evil"), + _valid_claims(kind="access"), + _valid_claims(user_id=""), + _valid_claims(nbf=0), + {k: v for k, v in _valid_claims().items() if k != "client_id"}, + {k: v for k, v in _valid_claims().items() if k != "exp"}, + ], +) +def test_signed_but_malformed_claims_are_rejected_without_raising(claims): + token = _sign_claims(claims) + assert isinstance(open_session_token(token, KEYS, NOW), SessionMalformed) + + +def test_signed_claims_with_exact_shape_open(): + token = _sign_claims(_valid_claims()) + opened = open_session_token(token, KEYS, NOW) + assert isinstance(opened, OpenedSessionToken) + assert opened.principal.user_id == "user-123" + + +def test_oversized_client_id_fails_mint_with_typed_error_not_truncation(): + principal = SessionPrincipal(user_id="user-123", client_id="c" * (MAX_SESSION_TOKEN_BYTES + 100)) + minted = mint_session_token(principal, KEYS, NOW) + assert isinstance(minted, SessionTokenTooLarge) + assert minted.max_bytes == MAX_SESSION_TOKEN_BYTES + + +def test_empty_principal_fields_rejected_at_construction(): + with pytest.raises(ValidationError): + SessionPrincipal(user_id="", client_id="c") + with pytest.raises(ValidationError): + SessionPrincipal(user_id="u", client_id="") + + +def test_short_signing_key_rejected_at_construction(): + with pytest.raises(ValidationError): + SessionKeys(signing_key=SecretStr("short")) + + +def test_two_mints_of_the_same_principal_are_distinct_tokens(): + first = mint_session_token(PRINCIPAL, KEYS, NOW) + second = mint_session_token(PRINCIPAL, KEYS, NOW) + assert isinstance(first, MintedSessionToken) and isinstance(second, MintedSessionToken) + assert first.token.get_secret_value() != second.token.get_secret_value() + + +def test_minted_token_repr_never_leaks_value(): + minted = mint_session_token(PRINCIPAL, KEYS, NOW) + assert isinstance(minted, MintedSessionToken) + assert minted.token.get_secret_value() not in repr(minted) 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 491fa023031..03f91260955 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 @@ -1759,6 +1759,42 @@ class TestMCPServerManager: assert client._resolved_auth is not None assert "authorization" not in {k.lower() for k in (client.extra_headers or {})} + @pytest.mark.asyncio + async def test_injected_authorization_does_not_shadow_m2m_minted_token(self): + """The M2M twin of the OBO shadow test: a guardrail/static Authorization must not displace + the gateway-minted client_credentials bearer. Dropping the resolved auth here would also + drop the one-shot 401 refetch that rides on it, so the resolver-owned credential is + authoritative exactly as for token_exchange and authorization_code.""" + from litellm.proxy._experimental.mcp_server.outbound_credentials.httpx_auth import ( + StaticHeaderAuth, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + + class _FakeProvider: + async def resolve_credentials(self, subject, server): + return Ok(StaticHeaderAuth("Bearer MINTED-M2M", header_name="Authorization")) + + manager = MCPServerManager(cred_provider=_FakeProvider()) + server = MCPServer( + server_id="m2m-shadow", + name="m2m-shadow-server", + url="https://up.example.com/mcp", + transport=MCPTransport.http, + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + client_id="cid", + client_secret="csec", + token_url="https://idp.example.com/token", + ) + + client = await manager._create_mcp_client( + server, + extra_headers={"Authorization": "Bearer signer-jwt"}, # simulate the JWT signer + ) + + assert client._resolved_auth is not None + assert "authorization" not in {k.lower() for k in (client.extra_headers or {})} + @pytest.mark.asyncio async def test_preflight_token_exchange_challenges_on_rejected_subject(self): """A subject the IdP rejects must raise the RFC 9728 401 challenge from the preflight, so a @@ -7760,24 +7796,59 @@ class TestCreateMcpClientV2Graft: assert client._resolved_auth.header_name == "Authorization" assert client._resolved_auth._header_value.get_secret_value() == f"Basic {encoded}" - async def test_m2m_client_credentials_defers_to_v1(self): - # M2M (oauth2 client_credentials) is not migrated: to_server_spec returns - # None, so the graft sets no resolved auth and leaves v1 in charge (v1 - # performs the client_credentials grant itself - the static - # authentication_token is never consumed for oauth2, so it does not flow - # to _mcp_auth_value). Per-user oauth2 (authorization_code) is migrated to - # v2 and is exercised separately. + async def test_m2m_client_credentials_resolves_via_v2(self): + # M2M (oauth2 client_credentials) is migrated: to_server_spec owns the server and the + # v2 arm mints the token through the injected source; nothing flows to v1's auth_value. + from litellm.proxy._experimental.mcp_server.outbound_credentials import ( + UpstreamCredentialProvider, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.client_credentials import ( + ClientCredentialsBearerAuth, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.oauth_token_store import ( + OAuthToken, + ) + from litellm.proxy._experimental.mcp_server.outbound_credentials.result import Ok + + class _FakeM2MSource: + async def get(self, server_id, config): + return Ok(OAuthToken(access_token="m2m-at")) + + async def refetch(self, server_id, config, failed_access_token): + return None + client = await MCPServerManager()._create_mcp_client( self._http_server( auth_type=MCPAuth.oauth2, oauth2_flow="client_credentials", - authentication_token="legacy-token", - ) + client_id="cid", + client_secret="csec", + token_url="https://idp.example.com/token", + ), + cred_provider=UpstreamCredentialProvider(client_credentials_source=_FakeM2MSource()), ) - assert client._resolved_auth is None + assert isinstance(client._resolved_auth, ClientCredentialsBearerAuth) + assert client._resolved_auth._access_token.get_secret_value() == "m2m-at" assert client._mcp_auth_value is None + async def test_m2m_client_credentials_incomplete_config_fails_closed(self): + # An M2M server missing its grant fields is still owned by v2 and surfaces a 500 + # misconfigured naming the missing fields, rather than deferring to v1 and connecting + # unauthenticated (which masked the upstream 401 as an empty tool list). + with pytest.raises(HTTPException) as exc_info: + await MCPServerManager()._create_mcp_client( + self._http_server( + auth_type=MCPAuth.oauth2, + oauth2_flow="client_credentials", + authentication_token="legacy-token", + ) + ) + + assert exc_info.value.status_code == 500 + assert "misconfigured" in str(exc_info.value.detail) + assert "token_url" in str(exc_info.value.detail) + async def test_static_token_missing_defers_to_v1(self): client = await MCPServerManager()._create_mcp_client( self._http_server(auth_type=MCPAuth.api_key, authentication_token=None) diff --git a/tests/test_litellm/proxy/auth/test_auth_utils.py b/tests/test_litellm/proxy/auth/test_auth_utils.py index b5d8727f7e6..72bd215b9be 100644 --- a/tests/test_litellm/proxy/auth/test_auth_utils.py +++ b/tests/test_litellm/proxy/auth/test_auth_utils.py @@ -1944,6 +1944,91 @@ class TestIsRequestBodySafeBlocksRivaUseSsl: ) +class TestIsRequestBodySafeBlocksBedrockTags: + """``bedrock_tags`` lands as AWS resource tags on Bedrock batch jobs + created with the proxy's AWS identity, so a caller-supplied value can + forge ownership or cost-allocation labels; like + ``aws_bedrock_project_id`` it is blocked without an admin opt-in.""" + + def test_bedrock_tags_in_request_body_is_rejected(self): + with pytest.raises(ValueError, match="bedrock_tags"): + is_request_body_safe( + request_body={ + "model": "bedrock-batch-opus", + "bedrock_tags": [{"key": "application", "value": "genai-proxy"}], + }, + general_settings={}, + llm_router=None, + model="bedrock-batch-opus", + ) + + def test_admin_opt_in_proxy_wide_allows_bedrock_tags(self): + assert ( + is_request_body_safe( + request_body={ + "model": "bedrock-batch-opus", + "bedrock_tags": [{"key": "application", "value": "genai-proxy"}], + }, + general_settings={"allow_client_side_credentials": True}, + llm_router=None, + model="bedrock-batch-opus", + ) + is True + ) + + def test_admin_opt_in_per_deployment_allows_bedrock_tags(self): + from litellm import Router + + router = Router( + model_list=[ + { + "model_name": "bedrock-batch-opus", + "litellm_params": { + "model": "bedrock/us.anthropic.claude-opus-4-7", + "configurable_clientside_auth_params": ["bedrock_tags"], + }, + } + ] + ) + assert ( + is_request_body_safe( + request_body={ + "model": "bedrock-batch-opus", + "bedrock_tags": [{"key": "application", "value": "genai-proxy"}], + }, + general_settings={}, + llm_router=router, + model="bedrock-batch-opus", + ) + is True + ) + + def test_per_deployment_opt_in_for_other_param_still_rejects_bedrock_tags(self): + from litellm import Router + + router = Router( + model_list=[ + { + "model_name": "bedrock-batch-opus", + "litellm_params": { + "model": "bedrock/us.anthropic.claude-opus-4-7", + "configurable_clientside_auth_params": ["api_base"], + }, + } + ] + ) + with pytest.raises(ValueError, match="bedrock_tags"): + is_request_body_safe( + request_body={ + "model": "bedrock-batch-opus", + "bedrock_tags": [{"key": "application", "value": "genai-proxy"}], + }, + general_settings={}, + llm_router=router, + model="bedrock-batch-opus", + ) + + # ── is_request_body_safe nested-config recursion (VERIA-6) ──────────────────── diff --git a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py index 79c5f3ea549..f3e5e2c9b71 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_model_management_endpoints.py @@ -1702,8 +1702,8 @@ class TestModelInfoEndpoint: async def test_model_info_accessible_model_success(self): """Test model_info returns model data for accessible models""" from litellm.proxy.proxy_server import model_info + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - # Mock user with access to specific models user_api_key_dict = UserAPIKeyAuth( user_id="test_user", api_key="test_key", @@ -1713,31 +1713,22 @@ class TestModelInfoEndpoint: with ( patch("litellm.proxy.proxy_server.llm_router") as mock_router, - patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, - patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, + patch("litellm.proxy.proxy_server.general_settings", {}), patch( - "litellm.proxy.proxy_server.get_complete_model_list" - ) as mock_get_complete_models, - patch("litellm.get_llm_provider") as mock_get_provider, + "litellm.proxy.utils.get_available_models_for_user", + new=AsyncMock(return_value=["gpt-4", "claude-3", "gpt-3.5-turbo"]), + ), + patch("litellm.get_llm_provider", return_value=(None, "openai", None, None)), ): - # Setup mocks - mock_router.get_model_names.return_value = [ - "gpt-4", - "claude-3", - "gpt-3.5-turbo", - ] - mock_router.get_model_access_groups.return_value = {} + mock_router.get_fully_blocked_model_names.return_value = set() + mock_router.get_model_list.return_value = [] mock_router.get_configured_token_limits.return_value = (None, None) - mock_get_key_models.return_value = ["gpt-4", "claude-3"] - mock_get_team_models.return_value = ["gpt-3.5-turbo"] - mock_get_complete_models.return_value = [ - "gpt-4", - "claude-3", - "gpt-3.5-turbo", - ] - mock_get_provider.return_value = (None, "openai", None, None) + mock_router.get_deployment_by_model_group_name.return_value = Deployment( + model_name="gpt-4", + litellm_params=LiteLLM_Params(model="openai/gpt-4"), + model_info=ModelInfo(id="gpt-4"), + ) - # Test accessible model result = await model_info( model_id="gpt-4", user_api_key_dict=user_api_key_dict ) @@ -1764,18 +1755,14 @@ class TestModelInfoEndpoint: with ( patch("litellm.proxy.proxy_server.llm_router") as mock_router, - patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, - patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, + patch("litellm.proxy.proxy_server.general_settings", {}), patch( - "litellm.proxy.proxy_server.get_complete_model_list" - ) as mock_get_complete_models, + "litellm.proxy.utils.get_available_models_for_user", + new=AsyncMock(return_value=["gpt-4"]), + ), ): - # Setup mocks - user only has access to gpt-4 - mock_router.get_model_names.return_value = ["gpt-4", "claude-3"] - mock_router.get_model_access_groups.return_value = {} - mock_get_key_models.return_value = ["gpt-4"] - mock_get_team_models.return_value = [] - mock_get_complete_models.return_value = ["gpt-4"] # Only gpt-4 accessible + mock_router.get_fully_blocked_model_names.return_value = set() + mock_router.get_model_list.return_value = [] # Test inaccessible model should raise 404 with pytest.raises(HTTPException) as exc_info: @@ -1791,8 +1778,8 @@ class TestModelInfoEndpoint: async def test_model_info_team_model_access(self): """Test model_info works with team model access""" from litellm.proxy.proxy_server import model_info + from litellm.types.router import Deployment, LiteLLM_Params, ModelInfo - # Mock user with team access user_api_key_dict = UserAPIKeyAuth( user_id="test_user", api_key="test_key", @@ -1803,23 +1790,22 @@ class TestModelInfoEndpoint: with ( patch("litellm.proxy.proxy_server.llm_router") as mock_router, - patch("litellm.proxy.proxy_server.get_key_models") as mock_get_key_models, - patch("litellm.proxy.proxy_server.get_team_models") as mock_get_team_models, + patch("litellm.proxy.proxy_server.general_settings", {}), patch( - "litellm.proxy.proxy_server.get_complete_model_list" - ) as mock_get_complete_models, - patch("litellm.get_llm_provider") as mock_get_provider, + "litellm.proxy.utils.get_available_models_for_user", + new=AsyncMock(return_value=["team-model-1"]), + ), + patch("litellm.get_llm_provider", return_value=(None, "custom", None, None)), ): - # Setup mocks - mock_router.get_model_names.return_value = ["team-model-1"] - mock_router.get_model_access_groups.return_value = {} + mock_router.get_fully_blocked_model_names.return_value = set() + mock_router.get_model_list.return_value = [] mock_router.get_configured_token_limits.return_value = (None, None) - mock_get_key_models.return_value = [] - mock_get_team_models.return_value = ["team-model-1"] - mock_get_complete_models.return_value = ["team-model-1"] - mock_get_provider.return_value = (None, "custom", None, None) + mock_router.get_deployment_by_model_group_name.return_value = Deployment( + model_name="team-model-1", + litellm_params=LiteLLM_Params(model="custom/team-model-1"), + model_info=ModelInfo(id="team-model-1"), + ) - # Test team model access result = await model_info( model_id="team-model-1", user_api_key_dict=user_api_key_dict ) @@ -2947,7 +2933,7 @@ class TestGetModelInfoWithIdBlocked: def test_get_model_info_with_id_propagates_blocked_true(self): from litellm.proxy.proxy_server import ProxyConfig - model = MagicMock() + model = MagicMock(spec=["model_id", "model_info", "blocked"]) model.model_id = "dep-1" model.model_info = {} model.blocked = True diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index bd8e92c3cc2..93b21c9d3c1 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -1244,7 +1244,8 @@ def test_ProxyConfig_get_model_info_with_id_returns_router_model_info(): assert snapshot == {"id": "m-1", "db_model": True, "blocked": False} -def test_ProxyConfig_get_model_info_with_id_missing_model_id_raises(): +def test_ProxyConfig_get_model_info_with_id_missing_model_id_raises(monkeypatch): + monkeypatch.setattr("litellm.proxy.proxy_server.premium_user", False) pc = ProxyConfig() # model with no model_id, no model_info — accessing .model_id will fail. bad = SimpleNamespace(model_info=None) 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 579f46c8c77..db72a7fb38c 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 @@ -1467,6 +1467,59 @@ async def test_ui_view_spend_logs_pagination(client, monkeypatch): app.dependency_overrides.pop(ps.user_api_key_auth, None) +@pytest.mark.parametrize( + "page_size, expected_status, expected_rows", + [ + (1000, 200, 1000), + (1001, 422, None), + ], +) +@pytest.mark.asyncio +async def test_ui_view_spend_logs_page_size_upper_bound( + client, monkeypatch, page_size, expected_status, expected_rows +): + mock_spend_logs = [ + { + "id": f"log{i}", + "request_id": f"req{i}", + "api_key": "sk-test-key", + "startTime": datetime.datetime.now(timezone.utc).isoformat(), + "model": "gpt-4", + } + for i in range(1200) + ] + + monkeypatch.setattr( + "litellm.proxy.proxy_server.prisma_client", + make_ui_spend_logs_mock_prisma(mock_spend_logs, lambda where: mock_spend_logs), + ) + + 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/v2", + params={ + "page": 1, + "page_size": page_size, + "start_date": start_date, + "end_date": end_date, + }, + headers={"Authorization": "Bearer sk-test"}, + ) + + assert response.status_code == expected_status + if expected_status == 200: + data = response.json() + assert data["page_size"] == page_size + assert len(data["data"]) == expected_rows + finally: + app.dependency_overrides.pop(ps.user_api_key_auth, None) + + @pytest.mark.asyncio async def test_ui_view_session_spend_logs_pagination(client, monkeypatch): mock_spend_logs = [ diff --git a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py index a1347aa111c..d60fff66c44 100644 --- a/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py +++ b/tests/test_litellm/responses/mcp/test_litellm_proxy_mcp_handler.py @@ -606,3 +606,45 @@ def test_completion_with_function_tools_works_without_fastapi_installed(): timeout=120, ) assert result.returncode == 0, result.stderr + + +def test_extract_tool_call_details_reads_anthropic_tool_use_input(): + """ + Regression test (LIT-4517): an Anthropic tool_use block carries its arguments + under `input`, not `arguments`. + + Given: A tool_use content block as /v1/messages returns it + When: The shared extractor reads it + Then: The arguments come back, so the MCP tool is called with them + + Reading only `arguments` fails silently rather than loudly: _parse_tool_arguments + turns the resulting None into {}, so the tool still executes, just with every + argument dropped. + """ + tool_use_block = { + "type": "tool_use", + "id": "toolu_01ABC", + "name": "read_wiki_structure", + "input": {"repoName": "BerriAI/litellm"}, + } + + name, arguments, call_id = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(tool_use_block) + + assert name == "read_wiki_structure" + assert call_id == "toolu_01ABC" + assert arguments == {"repoName": "BerriAI/litellm"} + assert LiteLLM_Proxy_MCP_Handler._parse_tool_arguments(arguments) == {"repoName": "BerriAI/litellm"} + + +def test_extract_tool_call_details_still_prefers_openai_arguments(): + """The OpenAI chat shape must keep winning; `input` is only the fallback.""" + openai_tool_call = { + "id": "call_123", + "function": {"name": "get_weather", "arguments": '{"city": "Paris"}'}, + } + + name, arguments, call_id = LiteLLM_Proxy_MCP_Handler._extract_tool_call_details(openai_tool_call) + + assert name == "get_weather" + assert call_id == "call_123" + assert arguments == '{"city": "Paris"}' diff --git a/tests/test_litellm/test_audio_transcription_rust_bridge.py b/tests/test_litellm/test_audio_transcription_rust_bridge.py new file mode 100644 index 00000000000..bbeb6c38f78 --- /dev/null +++ b/tests/test_litellm/test_audio_transcription_rust_bridge.py @@ -0,0 +1,151 @@ +import importlib + +import pytest + +import litellm +from litellm.llms.bedrock.audio_transcription import BedrockAudioTranscriptionRustDispatch + +rust_bridge = importlib.import_module("litellm.rust_bridge.transcription") + + +class SyncBridge: + def __init__(self) -> None: + self.calls: list[dict[str, object]] = [] + + def __call__( + self, + model: str, + audio: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: + self.calls.append({"model": model, "audio": audio, "optional_params": optional_params}) + return {"text": "hello"} + + +class AsyncBridge: + async def __call__( + self, + model: str, + audio: dict[str, object], + api_key: str | None, + api_base: str | None, + custom_llm_provider: str | None, + extra_headers: dict[str, object] | None, + optional_params: dict[str, object], + timeout_seconds: float | None, + ) -> dict[str, object]: + return {"text": "async"} + + +def test_enabled_sync_bridge_receives_audio() -> None: + bridge = SyncBridge() + rust_bridge.configure_rust_transcription(True, transcription=bridge) + result = rust_bridge.transcription( + model="mistral.voxtral-mini-3b-2507", + audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"}, + api_key=None, + api_base=None, + custom_llm_provider="bedrock", + extra_headers=None, + optional_params={"temperature": 0}, + timeout=5.0, + ) + assert result == {"text": "hello"} + assert bridge.calls[0]["audio"] == {"data": "AQI=", "format": "wav", "filename": "audio.wav"} + + +@pytest.mark.asyncio +async def test_enabled_async_bridge() -> None: + rust_bridge.configure_rust_transcription(True, atranscription=AsyncBridge()) + result = await rust_bridge.atranscription( + model="mistral.voxtral-mini-3b-2507", + audio={"data": "AQI=", "format": "wav", "filename": "audio.wav"}, + api_key=None, + api_base=None, + custom_llm_provider="bedrock", + extra_headers=None, + optional_params={}, + timeout=None, + ) + assert result == {"text": "async"} + + +def test_loader_returns_none_without_native_extension(monkeypatch: pytest.MonkeyPatch) -> None: + rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) + monkeypatch.setattr("litellm.rust_bridge.get_native_bridge", lambda: None) + assert rust_bridge.load_rust_transcription() is None + assert rust_bridge.load_rust_atranscription() is None + + +def test_dispatch_sync_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(rust_bridge, "transcription", lambda **_: None) + + with pytest.raises(RuntimeError, match="bridge is unavailable"): + BedrockAudioTranscriptionRustDispatch().audio_transcriptions( + model="bedrock/mistral.voxtral-mini-3b-2507", + audio_file=("audio.wav", b"audio", "audio/wav"), + api_key=None, + api_base=None, + custom_llm_provider="bedrock", + extra_headers=None, + optional_params={}, + timeout=5, + ) + + +@pytest.mark.asyncio +async def test_dispatch_async_path_requires_bridge(monkeypatch: pytest.MonkeyPatch) -> None: + async def unavailable(**_: object) -> None: + return None + + monkeypatch.setattr(rust_bridge, "atranscription", unavailable) + + with pytest.raises(RuntimeError, match="bridge is unavailable"): + await BedrockAudioTranscriptionRustDispatch().async_audio_transcriptions( + model="bedrock/mistral.voxtral-mini-3b-2507", + audio_file=("audio.wav", b"audio", "audio/wav"), + api_key=None, + api_base=None, + custom_llm_provider="bedrock", + extra_headers=None, + optional_params={}, + timeout=5, + ) + + +def test_bedrock_transcription_uses_rust_only_path() -> None: + rust_bridge.configure_rust_transcription( + transcription=lambda **_: {"text": "rust"}, + atranscription=None, + ) + try: + response = litellm.transcription( + model="bedrock/mistral.voxtral-mini-3b-2507", + file=("audio.wav", b"audio", "audio/wav"), + ) + finally: + rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) + + assert response.text == "rust" + + +@pytest.mark.asyncio +async def test_bedrock_atranscription_uses_rust_only_path() -> None: + async def rust_response(**_: object) -> dict[str, object]: + return {"text": "rust"} + + rust_bridge.configure_rust_transcription(transcription=None, atranscription=rust_response) + try: + response = await litellm.atranscription( + model="bedrock/mistral.voxtral-mini-3b-2507", + file=("audio.wav", b"audio", "audio/wav"), + ) + finally: + rust_bridge.configure_rust_transcription(transcription=None, atranscription=None) + + assert response.text == "rust" diff --git a/tests/test_litellm/test_router.py b/tests/test_litellm/test_router.py index 1c175bf6f44..cffc3dd0aba 100644 --- a/tests/test_litellm/test_router.py +++ b/tests/test_litellm/test_router.py @@ -5806,3 +5806,63 @@ def test_get_configured_token_limits_coerces_numeric_strings(): ) assert router.get_configured_token_limits("quoted-limits-model") == (32000, 8000) + + +@pytest.mark.asyncio +async def test_acreate_batch_request_bedrock_tags_override_deployment_tags(): + import httpx + + from litellm.llms.bedrock.common_utils import CommonBatchFilesUtils + + deployment_tags = [{"key": "application", "value": "config-level"}] + request_tags = [{"key": "application", "value": "request-level"}] + router = litellm.Router( + model_list=[ + { + "model_name": "bedrock-batch-model", + "litellm_params": { + "model": "bedrock/us.anthropic.claude-sonnet-5", + "aws_batch_role_arn": "arn:aws:iam::123:role/batch-role", + "aws_region_name": "us-west-2", + "bedrock_tags": deployment_tags, + }, + } + ] + ) + + def fake_response(): + return httpx.Response( + status_code=200, + json={ + "jobArn": "arn:aws:bedrock:us-west-2:123:model-invocation-job/abc1234567", + "status": "Submitted", + }, + ) + + mock_client = MagicMock() + mock_client.post = AsyncMock(side_effect=lambda *args, **kwargs: fake_response()) + + with patch.object( + CommonBatchFilesUtils, + "sign_aws_request", + return_value=({"Authorization": "signed"}, b"{}"), + ) as mock_sign, patch( + "litellm.llms.custom_httpx.llm_http_handler.get_async_httpx_client", + return_value=mock_client, + ): + await router.acreate_batch( + model="bedrock-batch-model", + input_file_id="s3://bucket/input.jsonl", + endpoint="/v1/chat/completions", + completion_window="24h", + ) + assert mock_sign.call_args.kwargs["data"]["tags"] == deployment_tags + + await router.acreate_batch( + model="bedrock-batch-model", + input_file_id="s3://bucket/input.jsonl", + endpoint="/v1/chat/completions", + completion_window="24h", + bedrock_tags=request_tags, + ) + assert mock_sign.call_args.kwargs["data"]["tags"] == request_tags diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutorouterTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutorouterTab.tsx new file mode 100644 index 00000000000..3474036528a --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/AutorouterTab.tsx @@ -0,0 +1,28 @@ +"use client"; + +import React from "react"; +import { Form } from "antd"; + +import AddAutoRouterTab from "@/components/add_model/add_auto_router_tab"; + +interface AutorouterTabProps { + accessToken: string | null; + userId: string | null; + userRole: string; +} + +const AutorouterTab: React.FC = ({ accessToken, userRole }) => { + const [form] = Form.useForm(); + + if (!accessToken) { + return null; + } + + return ( +
+ form.resetFields()} accessToken={accessToken} userRole={userRole} /> +
+ ); +}; + +export default AutorouterTab; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx index 92426d3ce04..46aa23fcfc0 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.test.tsx @@ -1,109 +1,34 @@ -import { render } from "@testing-library/react"; +import { fireEvent, render } from "@testing-library/react"; import { describe, expect, it, vi } from "vitest"; -import type { DailyData, SpendMetrics } from "@/components/UsagePage/types"; - -const mockUsePaginatedDailyActivity = vi.fn(); - -vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", () => ({ - usePaginatedDailyActivity: (args: unknown) => mockUsePaginatedDailyActivity(args), -})); - -vi.mock("@/components/networking", () => ({ - userDailyActivityCall: vi.fn(), -})); - -vi.mock("@/components/shared/advanced_date_picker", () => ({ - __esModule: true, - default: () =>
, -})); - -vi.mock("@/components/shared/charts", () => ({ - AreaChart: ({ data, categories }: { data: unknown; categories: string[] }) => ( -
- ), - DonutChart: ({ data, label }: { data: unknown; label: string }) => ( -
- ), -})); +vi.mock("./UsageTab", () => ({ __esModule: true, default: () =>
})); +vi.mock("./PromptCompressionTab", () => ({ __esModule: true, default: () =>
})); +vi.mock("./AutorouterTab", () => ({ __esModule: true, default: () =>
})); +vi.mock("./PromptCachingTab", () => ({ __esModule: true, default: () =>
})); import CostOptimizationView from "./CostOptimizationView"; -const baseMetrics = (overrides: Partial): SpendMetrics => ({ - spend: 0, - prompt_tokens: 0, - completion_tokens: 0, - total_tokens: 0, - api_requests: 0, - successful_requests: 0, - failed_requests: 0, - cache_read_input_tokens: 0, - cache_creation_input_tokens: 0, - ...overrides, -}); - -const day = (date: string, metrics: Partial): DailyData => ({ - date, - metrics: baseMetrics(metrics), - breakdown: { - models: {}, - model_groups: {}, - mcp_servers: {}, - providers: {}, - api_keys: {}, - entities: {}, - }, -}); - -const renderWith = (results: DailyData[]) => { - mockUsePaginatedDailyActivity.mockReturnValue({ data: { results }, loading: false, isFetchingMore: false }); - return render(); -}; +const renderView = () => render(); describe("CostOptimizationView", () => { - it("sums compression and caching dollars across days into the summary cards", () => { - const { getByText } = renderWith([ - day("2026-07-12", { - compression_savings_spend: 0.04, - prompt_caching_savings_spend: 0.006, - compression_saved_tokens: 40000, - }), - day("2026-07-13", { - compression_savings_spend: 0.1, - prompt_caching_savings_spend: 0.01, - compression_saved_tokens: 100000, - }), - ]); + it("renders all four cost-optimization tabs", () => { + const { getByText } = renderView(); - // compression 0.14 + caching 0.016 = 0.156 - expect(getByText("$0.1560")).toBeInTheDocument(); - expect(getByText("$0.1400")).toBeInTheDocument(); - expect(getByText("$0.0160")).toBeInTheDocument(); - expect(getByText("140,000 tokens compressed")).toBeInTheDocument(); + expect(getByText("Usage")).toBeInTheDocument(); + expect(getByText("Prompt Compression")).toBeInTheDocument(); + expect(getByText("Autorouter")).toBeInTheDocument(); + expect(getByText("Prompt Caching")).toBeInTheDocument(); }); - it("builds a per-day time series and per-driver donut from the daily rows", () => { - const { getByTestId } = renderWith([ - day("2026-07-12", { compression_savings_spend: 0.04, prompt_caching_savings_spend: 0.006 }), - day("2026-07-13", { compression_savings_spend: 0.1, prompt_caching_savings_spend: 0.01 }), - ]); + it("defaults to the Usage tab and switches the active tab on click", () => { + const { getByRole } = renderView(); - const series = JSON.parse(getByTestId("area-chart").getAttribute("data-series") ?? "[]"); - expect(series).toHaveLength(2); - expect(series[0]).toMatchObject({ Compression: 0.04, "Prompt caching": 0.006 }); - expect(series[1]).toMatchObject({ Compression: 0.1, "Prompt caching": 0.01 }); + expect(getByRole("tab", { name: "Usage" })).toHaveAttribute("aria-selected", "true"); + expect(getByRole("tab", { name: "Prompt Compression" })).toHaveAttribute("aria-selected", "false"); - const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]"); - expect(slices).toEqual([ - { driver: "Compression", usd: expect.closeTo(0.14, 5) }, - { driver: "Prompt caching", usd: expect.closeTo(0.016, 5) }, - ]); - }); + fireEvent.click(getByRole("tab", { name: "Prompt Compression" })); - it("omits a driver slice when that driver has no savings", () => { - const { getByTestId } = renderWith([day("2026-07-12", { compression_savings_spend: 0.04 })]); - - const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]"); - expect(slices).toEqual([{ driver: "Compression", usd: expect.closeTo(0.04, 5) }]); + expect(getByRole("tab", { name: "Usage" })).toHaveAttribute("aria-selected", "false"); + expect(getByRole("tab", { name: "Prompt Compression" })).toHaveAttribute("aria-selected", "true"); }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx index e45e3e23f1d..6e6830b8451 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/CostOptimizationView.tsx @@ -1,16 +1,13 @@ "use client"; -import React, { useMemo, useState } from "react"; +import React from "react"; import { PiggyBank } from "lucide-react"; +import { Alert, Tabs } from "antd"; -import { AreaChart, DonutChart } from "@/components/shared/charts"; -import AdvancedDatePicker from "@/components/shared/advanced_date_picker"; -import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; -import { userDailyActivityCall } from "@/components/networking"; -import { DailyData, SpendMetrics } from "@/components/UsagePage/types"; -import { formatNumberWithCommas } from "@/utils/dataUtils"; -import { all_admin_roles } from "@/utils/roles"; -import { usePaginatedDailyActivity } from "@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity"; +import UsageTab from "./UsageTab"; +import PromptCompressionTab from "./PromptCompressionTab"; +import AutorouterTab from "./AutorouterTab"; +import PromptCachingTab from "./PromptCachingTab"; interface CostOptimizationViewProps { accessToken: string | null; @@ -18,138 +15,62 @@ interface CostOptimizationViewProps { userRole: string; } -type DateRange = { from?: Date; to?: Date }; - -const THIRTY_DAYS_MS = 30 * 24 * 60 * 60 * 1000; - -const usd = (value: number): string => { - const decimals = value > 0 && value < 1 ? 4 : 2; - return `$${formatNumberWithCommas(value, decimals)}`; -}; - -const shortDate = (iso: string): string => - new Date(`${iso}T00:00:00`).toLocaleDateString("en-US", { month: "short", day: "numeric" }); - -const compressionOf = (m: SpendMetrics): number => m.compression_savings_spend ?? 0; -const cachingOf = (m: SpendMetrics): number => m.prompt_caching_savings_spend ?? 0; -const savedTokensOf = (m: SpendMetrics): number => m.compression_saved_tokens ?? 0; - -const SummaryCard = ({ label, value, hint }: { label: string; value: string; hint?: string }) => ( - - - {label} - - -

{value}

- {hint &&

{hint}

} -
-
-); - const CostOptimizationView: React.FC = ({ accessToken, userId, userRole }) => { - const initialFrom = useMemo(() => new Date(new Date().getTime() - THIRTY_DAYS_MS), []); - const initialTo = useMemo(() => new Date(), []); - const [dateValue, setDateValue] = useState({ from: initialFrom, to: initialTo }); - - const startTime = dateValue.from ?? null; - const endTime = dateValue.to ?? null; - const isAdmin = all_admin_roles.includes(userRole); - const effectiveUserId = isAdmin ? null : userId; - - const { data, loading, isFetchingMore } = usePaginatedDailyActivity({ - fetchFn: userDailyActivityCall, - args: [accessToken, startTime, endTime, effectiveUserId], - enabled: !!accessToken && !!startTime && !!endTime, - }); - - const results = data.results as DailyData[]; - - const compressionTotal = useMemo(() => results.reduce((sum, d) => sum + compressionOf(d.metrics), 0), [results]); - const cachingTotal = useMemo(() => results.reduce((sum, d) => sum + cachingOf(d.metrics), 0), [results]); - const savedTokensTotal = useMemo(() => results.reduce((sum, d) => sum + savedTokensOf(d.metrics), 0), [results]); - const totalSaved = compressionTotal + cachingTotal; - - const overTime = useMemo( - () => - results.map((d) => ({ - date: shortDate(d.date), - Compression: compressionOf(d.metrics), - "Prompt caching": cachingOf(d.metrics), - })), - [results], - ); - - const byDriver = useMemo( - () => - [ - { driver: "Compression", usd: compressionTotal }, - { driver: "Prompt caching", usd: cachingTotal }, - ].filter((d) => d.usd > 0), - [compressionTotal, cachingTotal], - ); + const items = [ + { + key: "usage", + label: "Usage", + children: , + }, + { + key: "compression", + label: "Prompt Compression", + children: , + }, + { + key: "autorouter", + label: "Autorouter", + children: , + }, + { + key: "caching", + label: "Prompt Caching", + children: , + }, + ]; return (
-
-
-
- -

Cost Optimization

-
-

- Money saved by prompt compression and prompt caching across your requests -

+
+
+ +

Cost Optimization

- setDateValue(v)} /> +

+ Track and configure the mechanisms that save you money: prompt compression, prompt caching, and auto routing +

-
- - - -
+ + Have feedback? Join the discussion{" "} + + here + + + } + /> -
- - - Savings over time - - - - - - - - Savings by driver - - - - - -
+
); }; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.tsx new file mode 100644 index 00000000000..e6f73088824 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCachingTab.tsx @@ -0,0 +1,52 @@ +"use client"; + +import React, { useCallback, useEffect, useState } from "react"; + +import { getGeneralSettingsCall } from "@/components/networking"; +import NotificationsManager from "@/components/molecules/notifications_manager"; +import { + PromptCachingPanel, + generalSettingsItem, +} from "@/app/(dashboard)/router-settings/_components/general_settings"; + +interface PromptCachingTabProps { + accessToken: string | null; +} + +const PromptCachingTab: React.FC = ({ accessToken }) => { + const [settings, setSettings] = useState([]); + + const loadSettings = useCallback(() => { + if (!accessToken) { + return; + } + getGeneralSettingsCall(accessToken) + .then((data: generalSettingsItem[]) => setSettings(data)) + .catch((error) => { + console.error("Failed to load prompt caching settings:", error); + NotificationsManager.fromBackend("Failed to load prompt caching settings"); + }); + }, [accessToken]); + + useEffect(() => { + loadSettings(); + }, [loadSettings]); + + const handleChange = (fieldName: string, newValue: unknown) => { + setSettings((prev) => + prev.map((setting) => (setting.field_name === fieldName ? { ...setting, field_value: newValue } : setting)), + ); + }; + + if (!accessToken) { + return null; + } + + return ( +
+ +
+ ); +}; + +export default PromptCachingTab; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCompressionTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCompressionTab.tsx new file mode 100644 index 00000000000..071ad1d6bd5 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/PromptCompressionTab.tsx @@ -0,0 +1,176 @@ +"use client"; + +import React, { useCallback, useEffect, useState } from "react"; +import { Button, Form, Input, Switch } from "antd"; + +import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { createGuardrailCall, getGuardrailsList } from "@/components/networking"; +import NotificationsManager from "@/components/molecules/notifications_manager"; +import { + buildCompressionGuardrailPayload, + compressionGuardrailsOf, + GuardrailListItem, + GuardrailListResponse, +} from "./helpers"; + +interface PromptCompressionTabProps { + accessToken: string | null; +} + +interface CompressionFormValues { + name: string; + apiBase: string; + defaultOn: boolean; +} + +const PromptCompressionTab: React.FC = ({ accessToken }) => { + const [form] = Form.useForm(); + const [guardrails, setGuardrails] = useState([]); + const [isLoading, setIsLoading] = useState(true); + const [isSaving, setIsSaving] = useState(false); + + const loadGuardrails = useCallback(() => { + if (!accessToken) { + return; + } + getGuardrailsList(accessToken) + .then((response) => setGuardrails(compressionGuardrailsOf(response as GuardrailListResponse))) + .catch((error) => { + console.error("Failed to load compression guardrails:", error); + NotificationsManager.fromBackend("Failed to load compression guardrails"); + }) + .finally(() => setIsLoading(false)); + }, [accessToken]); + + useEffect(() => { + loadGuardrails(); + }, [loadGuardrails]); + + const handleAdd = async (values: CompressionFormValues) => { + if (!accessToken) { + return; + } + setIsSaving(true); + try { + await createGuardrailCall( + accessToken, + buildCompressionGuardrailPayload({ + name: values.name, + apiBase: values.apiBase, + defaultOn: values.defaultOn ?? true, + }), + ); + NotificationsManager.success("Compression guardrail created"); + form.resetFields(); + await loadGuardrails(); + } catch (error) { + console.error("Failed to create compression guardrail:", error); + NotificationsManager.fromBackend("Failed to create compression guardrail"); + } finally { + setIsSaving(false); + } + }; + + return ( +
+ + + Headroom prompt compression + + +

+ Headroom is a native LiteLLM guardrail that compresses your prompts before they reach the model, so you pay + for fewer input tokens. The tokens it removes are priced and shown on the Usage tab as compression savings.{" "} + + Headroom setup docs + +

+ {isLoading &&

Loading...

} + {!isLoading && guardrails.length === 0 && ( +

+ No prompt compression guardrails configured yet. Add one below to start saving on input tokens +

+ )} + {!isLoading && guardrails.length > 0 && ( +
    + {guardrails.map((guardrail) => ( +
  • +
    +

    {guardrail.guardrail_name}

    +

    {guardrail.litellm_params?.api_base ?? ""}

    +
    + + {guardrail.litellm_params?.default_on ? "Always on" : "Opt-in"} + +
  • + ))} +
+ )} +
+
+ + + + Add Headroom compression guardrail + + +
+ + + + + + + + + +
+

+ Applying compression to all requests is available to all users. Enabling it selectively per key or team + is a LiteLLM Enterprise feature. Get a trial key{" "} + + here + +

+
+
+ +
+
+
+
+
+ ); +}; + +export default PromptCompressionTab; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx new file mode 100644 index 00000000000..048d9a34d31 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.test.tsx @@ -0,0 +1,108 @@ +import { render } from "@testing-library/react"; +import { describe, expect, it, vi } from "vitest"; + +import type { DailyData, SpendMetrics } from "@/components/UsagePage/types"; + +const mockUsePaginatedDailyActivity = vi.fn(); + +vi.mock("@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity", () => ({ + usePaginatedDailyActivity: (args: unknown) => mockUsePaginatedDailyActivity(args), +})); + +vi.mock("@/components/networking", () => ({ + userDailyActivityCall: vi.fn(), +})); + +vi.mock("@/components/shared/advanced_date_picker", () => ({ + __esModule: true, + default: () =>
, +})); + +vi.mock("@/components/shared/charts", () => ({ + AreaChart: ({ data, categories }: { data: unknown; categories: string[] }) => ( +
+ ), + DonutChart: ({ data, label }: { data: unknown; label: string }) => ( +
+ ), +})); + +import UsageTab from "./UsageTab"; + +const baseMetrics = (overrides: Partial): SpendMetrics => ({ + spend: 0, + prompt_tokens: 0, + completion_tokens: 0, + total_tokens: 0, + api_requests: 0, + successful_requests: 0, + failed_requests: 0, + cache_read_input_tokens: 0, + cache_creation_input_tokens: 0, + ...overrides, +}); + +const day = (date: string, metrics: Partial): DailyData => ({ + date, + metrics: baseMetrics(metrics), + breakdown: { + models: {}, + model_groups: {}, + mcp_servers: {}, + providers: {}, + api_keys: {}, + entities: {}, + }, +}); + +const renderWith = (results: DailyData[]) => { + mockUsePaginatedDailyActivity.mockReturnValue({ data: { results }, loading: false, isFetchingMore: false }); + return render(); +}; + +describe("UsageTab", () => { + it("sums compression and caching dollars across days into the summary cards", () => { + const { getByText } = renderWith([ + day("2026-07-12", { + compression_savings_spend: 0.04, + prompt_caching_savings_spend: 0.006, + compression_saved_tokens: 40000, + }), + day("2026-07-13", { + compression_savings_spend: 0.1, + prompt_caching_savings_spend: 0.01, + compression_saved_tokens: 100000, + }), + ]); + + expect(getByText("$0.1560")).toBeInTheDocument(); + expect(getByText("$0.1400")).toBeInTheDocument(); + expect(getByText("$0.0160")).toBeInTheDocument(); + expect(getByText("140,000 tokens compressed")).toBeInTheDocument(); + }); + + it("builds a per-day time series and per-driver donut from the daily rows", () => { + const { getByTestId } = renderWith([ + day("2026-07-12", { compression_savings_spend: 0.04, prompt_caching_savings_spend: 0.006 }), + day("2026-07-13", { compression_savings_spend: 0.1, prompt_caching_savings_spend: 0.01 }), + ]); + + const series = JSON.parse(getByTestId("area-chart").getAttribute("data-series") ?? "[]"); + expect(series).toHaveLength(2); + expect(series[0]).toMatchObject({ Compression: 0.04, "Prompt caching": 0.006 }); + expect(series[1]).toMatchObject({ Compression: 0.1, "Prompt caching": 0.01 }); + + const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]"); + expect(slices).toEqual([ + { driver: "Compression", usd: expect.closeTo(0.14, 5) }, + { driver: "Prompt caching", usd: expect.closeTo(0.016, 5) }, + ]); + }); + + it("omits a driver slice when that driver has no savings", () => { + const { getByTestId } = renderWith([day("2026-07-12", { compression_savings_spend: 0.04 })]); + + const slices = JSON.parse(getByTestId("donut-chart").getAttribute("data-slices") ?? "[]"); + expect(slices).toEqual([{ driver: "Compression", usd: expect.closeTo(0.04, 5) }]); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.tsx b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.tsx new file mode 100644 index 00000000000..8e6fc40b5ad --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/UsageTab.tsx @@ -0,0 +1,184 @@ +"use client"; + +import React, { useMemo, useState } from "react"; +import { Collapse } from "antd"; + +import { AreaChart, DonutChart } from "@/components/shared/charts"; +import AdvancedDatePicker from "@/components/shared/advanced_date_picker"; +import { Card, CardContent, CardHeader, CardTitle } from "@/components/ui/card"; +import { userDailyActivityCall } from "@/components/networking"; +import { DailyData, SpendMetrics } from "@/components/UsagePage/types"; +import { formatNumberWithCommas } from "@/utils/dataUtils"; +import { all_admin_roles } from "@/utils/roles"; +import { usePaginatedDailyActivity } from "@/app/(dashboard)/usage/_components/hooks/usePaginatedDailyActivity"; + +interface UsageTabProps { + accessToken: string | null; + userId: string | null; + userRole: string; +} + +type DateRange = { from?: Date; to?: Date }; + +const THIRTY_DAYS_MS = 30 * 24 * 60 * 60 * 1000; + +const usd = (value: number): string => { + const decimals = value > 0 && value < 1 ? 4 : 2; + return `$${formatNumberWithCommas(value, decimals)}`; +}; + +const shortDate = (iso: string): string => + new Date(`${iso}T00:00:00`).toLocaleDateString("en-US", { month: "short", day: "numeric" }); + +const compressionOf = (m: SpendMetrics): number => m.compression_savings_spend ?? 0; +const cachingOf = (m: SpendMetrics): number => m.prompt_caching_savings_spend ?? 0; +const savedTokensOf = (m: SpendMetrics): number => m.compression_saved_tokens ?? 0; + +const MethodologyNote = () => ( + How savings are calculated, + children: ( +
+

+ Savings are computed for each request when it is logged, using the provider's reported usage and the + model's pricing, then summed into a daily rollup. Totals below are read from that rollup over the + selected date range, so the numbers never require a scan of raw request logs. +

+

+ Compression savings are the tokens Headroom removed before the call, priced at the model's input + rate: compression_saved_tokens * input_cost_per_token +

+

+ Prompt caching savings are the tokens the provider served from cache (Anthropic{" "} + cache_read_input_tokens, or OpenAI-style prompt_tokens_details.cached_tokens), + priced at the discount between the normal input rate and the cache-read rate:{" "} + cache_read_input_tokens * max(input_cost_per_token - cache_read_input_token_cost, 0) +

+

+ Total saved is the sum of both drivers. Models without a separate cache-read price in the pricing map + contribute zero caching savings rather than erroring. +

+
+ ), + }, + ]} + /> +); + +const SummaryCard = ({ label, value, hint }: { label: string; value: string; hint?: string }) => ( + + + {label} + + +

{value}

+ {hint &&

{hint}

} +
+
+); + +const UsageTab: React.FC = ({ accessToken, userId, userRole }) => { + const initialFrom = useMemo(() => new Date(new Date().getTime() - THIRTY_DAYS_MS), []); + const initialTo = useMemo(() => new Date(), []); + const [dateValue, setDateValue] = useState({ from: initialFrom, to: initialTo }); + + const startTime = dateValue.from ?? null; + const endTime = dateValue.to ?? null; + const isAdmin = all_admin_roles.includes(userRole); + const effectiveUserId = isAdmin ? null : userId; + + const { data, loading, isFetchingMore } = usePaginatedDailyActivity({ + fetchFn: userDailyActivityCall, + args: [accessToken, startTime, endTime, effectiveUserId], + enabled: !!accessToken && !!startTime && !!endTime, + }); + + const results = data.results as DailyData[]; + + const compressionTotal = useMemo(() => results.reduce((sum, d) => sum + compressionOf(d.metrics), 0), [results]); + const cachingTotal = useMemo(() => results.reduce((sum, d) => sum + cachingOf(d.metrics), 0), [results]); + const savedTokensTotal = useMemo(() => results.reduce((sum, d) => sum + savedTokensOf(d.metrics), 0), [results]); + const totalSaved = compressionTotal + cachingTotal; + + const overTime = useMemo( + () => + results.map((d) => ({ + date: shortDate(d.date), + Compression: compressionOf(d.metrics), + "Prompt caching": cachingOf(d.metrics), + })), + [results], + ); + + const byDriver = useMemo( + () => + [ + { driver: "Compression", usd: compressionTotal }, + { driver: "Prompt caching", usd: cachingTotal }, + ].filter((d) => d.usd > 0), + [compressionTotal, cachingTotal], + ); + + return ( +
+
+ + setDateValue(v)} /> +
+ +
+ + + +
+ +
+ + + Savings over time + + + + + + + + Savings by driver + + + + + +
+
+ ); +}; + +export default UsageTab; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/helpers.test.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/helpers.test.ts new file mode 100644 index 00000000000..239cded0815 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/helpers.test.ts @@ -0,0 +1,37 @@ +import { describe, expect, it } from "vitest"; + +import { buildCompressionGuardrailPayload, compressionGuardrailsOf } from "./helpers"; + +describe("compressionGuardrailsOf", () => { + it("keeps only headroom-provider guardrails and drops others", () => { + const filtered = compressionGuardrailsOf({ + guardrails: [ + { guardrail_id: "1", guardrail_name: "headroom-compression", litellm_params: { guardrail: "headroom" } }, + { guardrail_id: "2", guardrail_name: "pii-masker", litellm_params: { guardrail: "presidio" } }, + { guardrail_id: "3", guardrail_name: "no-params", litellm_params: null }, + ], + }); + + expect(filtered.map((g) => g.guardrail_id)).toEqual(["1"]); + }); +}); + +describe("buildCompressionGuardrailPayload", () => { + it("builds a headroom guardrail payload with trimmed fields", () => { + const payload = buildCompressionGuardrailPayload({ + name: " headroom-compression ", + apiBase: " https://compress ", + defaultOn: false, + }); + + expect(payload).toEqual({ + guardrail_name: "headroom-compression", + litellm_params: { + guardrail: "headroom", + mode: "pre_call", + api_base: "https://compress", + default_on: false, + }, + }); + }); +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/helpers.ts b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/helpers.ts new file mode 100644 index 00000000000..7c8c92f8890 --- /dev/null +++ b/ui/litellm-dashboard/src/app/(dashboard)/cost-optimization/_components/helpers.ts @@ -0,0 +1,39 @@ +export interface GuardrailLitellmParams { + guardrail?: string | null; + api_base?: string | null; + default_on?: boolean | null; +} + +export interface GuardrailListItem { + guardrail_id: string; + guardrail_name: string | null; + litellm_params?: GuardrailLitellmParams | null; +} + +export interface GuardrailListResponse { + guardrails?: GuardrailListItem[]; +} + +export const COMPRESSION_GUARDRAIL_PROVIDER = "headroom"; + +export const isCompressionGuardrail = (guardrail: GuardrailListItem): boolean => + (guardrail.litellm_params?.guardrail ?? "").toLowerCase() === COMPRESSION_GUARDRAIL_PROVIDER; + +export const compressionGuardrailsOf = (response: GuardrailListResponse): GuardrailListItem[] => + (response.guardrails ?? []).filter(isCompressionGuardrail); + +export interface CompressionGuardrailInput { + name: string; + apiBase: string; + defaultOn: boolean; +} + +export const buildCompressionGuardrailPayload = (input: CompressionGuardrailInput): Record => ({ + guardrail_name: input.name.trim(), + litellm_params: { + guardrail: COMPRESSION_GUARDRAIL_PROVIDER, + mode: "pre_call", + api_base: input.apiBase.trim(), + default_on: input.defaultOn, + }, +}); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx index d8f927b9a63..d2cf27e0c8b 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/components/chat_ui/ChatUI.tsx @@ -98,7 +98,12 @@ interface ChatUIProps { fixedModel?: string; } -const MCP_SUPPORTED_ENDPOINTS = new Set([EndpointType.CHAT, EndpointType.RESPONSES, EndpointType.MCP]); +const MCP_SUPPORTED_ENDPOINTS = new Set([ + EndpointType.CHAT, + EndpointType.RESPONSES, + EndpointType.MCP, + EndpointType.ANTHROPIC_MESSAGES, +]); const CUSTOM_MODEL_DEBOUNCE_WAIT_MS = 500; @@ -870,8 +875,11 @@ const ChatUI: React.FC = ({ selectedVectorStores.length > 0 ? selectedVectorStores : undefined, selectedGuardrails.length > 0 ? selectedGuardrails : undefined, selectedPolicies.length > 0 ? selectedPolicies : undefined, - selectedMCPServers, // Pass the selected tools array + selectedMCPServers, customProxyBaseUrl || undefined, + mcpServers, + mcpServerToolRestrictions, + mcpToolsets, ); } else if (endpointType === EndpointType.EMBEDDINGS) { await makeOpenAIEmbeddingsRequest( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx b/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx index ed2b4280b79..4319315396a 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/playground/llm_calls/anthropic_messages.tsx @@ -1,6 +1,8 @@ import Anthropic from "@anthropic-ai/sdk"; import { MessageType } from "@/components/chat_ui/types"; import { TokenUsage } from "@/components/chat_ui/ResponseMetrics"; +import { buildMcpToolBlocks } from "@/components/llm_calls/mcp_tool_blocks"; +import { MCPServer, MCPToolset } from "@/components/mcp_tools/types"; import { getProxyBaseUrl } from "@/components/networking"; import NotificationManager from "@/components/molecules/notifications_manager"; @@ -18,8 +20,11 @@ export async function makeAnthropicMessagesRequest( vector_store_ids?: string[], guardrails?: string[], policies?: string[], - selectedMCPTools?: string[], + selectedMCPServers?: string[], customBaseUrl?: string, + mcpServers?: MCPServer[], + mcpServerToolRestrictions?: Record, + mcpToolsets?: MCPToolset[], ) { if (!accessToken) { throw new Error("Virtual Key is required"); @@ -58,6 +63,13 @@ export async function makeAnthropicMessagesRequest( litellm_trace_id: traceId, }; + const tools = buildMcpToolBlocks({ + selectedMCPServers, + mcpServers, + mcpToolsets, + mcpServerToolRestrictions, + }); + if (tools.length > 0) requestBody.tools = tools; if (vector_store_ids) requestBody.vector_store_ids = vector_store_ids; if (guardrails) requestBody.guardrails = guardrails; if (policies) requestBody.policies = policies; diff --git a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx index 1e8658d5104..dda7a23a8d4 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/router-settings/_components/general_settings.tsx @@ -33,7 +33,7 @@ interface GeneralSettingsPageProps { userID: string | null; } -interface generalSettingsItem { +export interface generalSettingsItem { field_name: string; field_type: string; field_value: any; @@ -90,7 +90,7 @@ const SettingValueEditor: React.FC<{ return null; }; -const PromptCachingPanel: React.FC<{ +export const PromptCachingPanel: React.FC<{ accessToken: string; settings: generalSettingsItem[]; onChange: (fieldName: string, newValue: any) => void; diff --git a/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.test.ts b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.test.ts new file mode 100644 index 00000000000..62dd4d2631c --- /dev/null +++ b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.test.ts @@ -0,0 +1,63 @@ +import { describe, it, expect } from "vitest"; +import { buildMcpToolBlocks } from "./mcp_tool_blocks"; +import { MCPServer, MCPToolset } from "@/components/mcp_tools/types"; + +const server = (over: Partial): MCPServer => + ({ + server_id: "id-1", + server_name: "deepwiki", + alias: "wiki", + url: "", + transport: "http", + auth_type: "none", + ...over, + }) as any; + +describe("buildMcpToolBlocks", () => { + it("returns no blocks when nothing is selected", () => { + expect(buildMcpToolBlocks({ selectedMCPServers: [] })).toEqual([]); + expect(buildMcpToolBlocks({ selectedMCPServers: undefined })).toEqual([]); + }); + + it("routes by server_name, not alias, so colliding aliases cannot cross-route", () => { + const [block] = buildMcpToolBlocks({ + selectedMCPServers: ["id-1"], + mcpServers: [server({})], + }); + expect(block.server_url).toBe("litellm_proxy/mcp/deepwiki"); + expect(block.server_label).toBe("deepwiki"); + }); + + it("does not percent-encode the name; the gateway splits the raw path and never decodes", () => { + const [block] = buildMcpToolBlocks({ + selectedMCPServers: ["id-1"], + mcpServers: [server({ server_name: "my server" }) as any], + }); + expect(block.server_url).toBe("litellm_proxy/mcp/my server"); + expect(block.server_url).not.toContain("%20"); + }); + + it("passes per-server tool restrictions through as allowed_tools", () => { + const [block] = buildMcpToolBlocks({ + selectedMCPServers: ["id-1"], + mcpServers: [server({})], + mcpServerToolRestrictions: { "id-1": ["read_wiki_structure"] }, + }); + expect(block.allowed_tools).toEqual(["read_wiki_structure"]); + }); + + it("collapses the all-servers sentinel to a single proxy-wide block", () => { + expect(buildMcpToolBlocks({ selectedMCPServers: ["__all__", "id-1"] })).toEqual([ + { type: "mcp", server_label: "litellm", server_url: "litellm_proxy/mcp", require_approval: "never" }, + ]); + }); + + it("routes a toolset by its name", () => { + const toolset = { toolset_id: "ts-1", toolset_name: "docs" } as MCPToolset; + const [block] = buildMcpToolBlocks({ + selectedMCPServers: ["toolset:ts-1"], + mcpToolsets: [toolset], + }); + expect(block.server_url).toBe("litellm_proxy/mcp/docs"); + }); +}); diff --git a/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts new file mode 100644 index 00000000000..401d9fd9c84 --- /dev/null +++ b/ui/litellm-dashboard/src/components/llm_calls/mcp_tool_blocks.ts @@ -0,0 +1,83 @@ +import { MCPServer, MCPToolset } from "@/components/mcp_tools/types"; + +export const ALL_MCP_SERVERS_SENTINEL = "__all__"; +const TOOLSET_PREFIX = "toolset:"; + +export interface McpToolBlock { + type: "mcp"; + server_label: string; + server_url: string; + require_approval: "never"; + allowed_tools?: string[]; +} + +export interface BuildMcpToolBlocksArgs { + selectedMCPServers?: string[]; + mcpServers?: MCPServer[]; + mcpToolsets?: MCPToolset[]; + mcpServerToolRestrictions?: Record; +} + +/** + * Build the litellm_proxy MCP reference blocks for a playground request. + * + * Every endpoint that supports MCP sends the same reference shape; the gateway + * expands it server side and each endpoint's own transformation decides the + * final tool shape. Keeping one builder here stops the endpoints from drifting + * apart on routing name, label uniqueness, or escaping. + * + * server_name is used for both routing and labelling because it is the unique + * registered identifier; aliases can collide across servers, and a duplicated + * server_label causes silent tool-routing failures. + * + * The name is not percent-encoded: the gateway resolves it with a raw + * `server_url.split("/")[-1]` and never url-decodes, so an encoded name would + * fail server lookup rather than round-trip. + */ +export function buildMcpToolBlocks({ + selectedMCPServers, + mcpServers, + mcpToolsets, + mcpServerToolRestrictions, +}: BuildMcpToolBlocksArgs): McpToolBlock[] { + if (!selectedMCPServers || selectedMCPServers.length === 0) { + return []; + } + + if (selectedMCPServers.includes(ALL_MCP_SERVERS_SENTINEL)) { + return [ + { + type: "mcp", + server_label: "litellm", + server_url: "litellm_proxy/mcp", + require_approval: "never", + }, + ]; + } + + return selectedMCPServers.map((serverId) => { + if (serverId.startsWith(TOOLSET_PREFIX)) { + const toolsetId = serverId.slice(TOOLSET_PREFIX.length); + const toolset = mcpToolsets?.find((t) => t.toolset_id === toolsetId); + const toolsetName = toolset?.toolset_name || toolsetId; + return { + type: "mcp", + server_label: toolsetName, + server_url: `litellm_proxy/mcp/${toolsetName}`, + require_approval: "never", + }; + } + + const server = mcpServers?.find((s) => s.server_id === serverId); + const routeName = server?.server_name || serverId; + const allowedTools = mcpServerToolRestrictions?.[serverId] || []; + + return { + type: "mcp", + server_label: routeName, + server_url: `litellm_proxy/mcp/${routeName}`, + require_approval: "never", + ...(allowedTools.length > 0 ? { allowed_tools: allowedTools } : {}), + }; + }); +} diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx index a57676d07b8..dcc72ccac9c 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.test.tsx @@ -1,8 +1,8 @@ import * as networking from "@/components/networking"; -import { screen, waitFor } from "@testing-library/react"; +import { screen, waitFor, within } from "@testing-library/react"; import userEvent from "@testing-library/user-event"; import { afterEach, beforeEach, describe, expect, it, vi } from "vitest"; -import { renderWithProviders } from "../../../tests/test-utils"; +import { renderWithProviders, testQueryClient } from "../../../tests/test-utils"; import TeamInfoView from "./TeamInfo"; vi.mock("@/components/networking", () => ({ @@ -1024,4 +1024,99 @@ describe("TeamInfoView", () => { expect(payload).not.toHaveProperty("model_aliases"); }); }); + + describe("guardrails dropdown grouping", () => { + const guardrail = (name: string, defaultOn: boolean) => ({ + guardrail_name: name, + litellm_params: { default_on: defaultOn }, + }); + + const openGuardrailsDropdown = async (user: ReturnType) => { + renderWithProviders(); + + await waitFor(() => { + const teamNameElements = screen.queryAllByText("Test Team"); + expect(teamNameElements.length).toBeGreaterThan(0); + }); + + await user.click(screen.getByRole("tab", { name: "Settings" })); + + await waitFor(() => { + expect(screen.getByRole("button", { name: /edit settings/i })).toBeInTheDocument(); + }); + + await user.click(screen.getByRole("button", { name: /edit settings/i })); + + await waitFor(() => { + expect(screen.getByLabelText(/^Guardrails/)).toBeInTheDocument(); + }); + + const dropdownsBefore = new Set(document.querySelectorAll(".ant-select-dropdown")); + + await user.click(screen.getByLabelText(/^Guardrails/)); + + return waitFor( + () => { + const opened = Array.from(document.querySelectorAll(".ant-select-dropdown")).find( + (el) => !dropdownsBefore.has(el), + ); + expect(opened).toBeDefined(); + return opened as HTMLElement; + }, + { timeout: 5000 }, + ); + }; + + beforeEach(() => { + testQueryClient.clear(); + vi.mocked(networking.teamInfoCall).mockResolvedValue(createMockTeamData()); + }); + + it("should not render the Global or Other group headers when no global guardrails exist", async () => { + const user = userEvent.setup({ delay: null }); + vi.mocked(networking.getGuardrailsList).mockResolvedValue({ + guardrails: [guardrail("dwacxzcz", false), guardrail("dwadsa", false)], + }); + + const dropdown = await openGuardrailsDropdown(user); + + await waitFor(() => { + expect(within(dropdown).getByTitle("dwacxzcz")).toBeInTheDocument(); + }); + expect(within(dropdown).getByTitle("dwadsa")).toBeInTheDocument(); + expect(within(dropdown).queryByText("Global")).not.toBeInTheDocument(); + expect(within(dropdown).queryByText("Other")).not.toBeInTheDocument(); + }); + + it("should not render the Global or Other group headers when every guardrail is global", async () => { + const user = userEvent.setup({ delay: null }); + vi.mocked(networking.getGuardrailsList).mockResolvedValue({ + guardrails: [guardrail("always-on", true)], + }); + + const dropdown = await openGuardrailsDropdown(user); + + await waitFor(() => { + expect(within(dropdown).getByTitle("always-on")).toBeInTheDocument(); + }); + expect(within(dropdown).queryByText("Global")).not.toBeInTheDocument(); + expect(within(dropdown).queryByText("Other")).not.toBeInTheDocument(); + }); + + it("should render both group headers when global and non-global guardrails exist", async () => { + const user = userEvent.setup({ delay: null }); + vi.mocked(networking.getGuardrailsList).mockResolvedValue({ + guardrails: [guardrail("always-on", true), guardrail("opt-in", false)], + }); + + const dropdown = await openGuardrailsDropdown(user); + + await waitFor(() => { + expect(within(dropdown).getByText("Global")).toBeInTheDocument(); + }); + expect(within(dropdown).getByText("Other")).toBeInTheDocument(); + expect(within(dropdown).getByTitle("always-on")).toBeInTheDocument(); + expect(within(dropdown).getByTitle("opt-in")).toBeInTheDocument(); + }); + }); }); diff --git a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx index 74f201b40c6..eaf2faa08ee 100644 --- a/ui/litellm-dashboard/src/components/team/TeamInfo.tsx +++ b/ui/litellm-dashboard/src/components/team/TeamInfo.tsx @@ -14,7 +14,7 @@ import { teamMemberUpdateCall, teamUpdateCall, } from "@/components/networking"; -import { useGuardrails } from "@/app/(dashboard)/hooks/guardrails/useGuardrails"; +import { useGuardrails, GuardrailListItem } from "@/app/(dashboard)/hooks/guardrails/useGuardrails"; import { formatNumberWithCommas } from "@/utils/dataUtils"; import { mapEmptyStringToNull } from "@/utils/keyUpdateUtils"; import { isProxyAdminRole } from "@/utils/roles"; @@ -679,6 +679,16 @@ const TeamInfoView: React.FC = ({ ? nonGlobalOptIns : [...Array.from(globalGuardrailNames).filter((n) => !optedOutGlobals.has(n)), ...nonGlobalOptIns]; + const allGuardrails: GuardrailListItem[] = guardrailsData?.guardrails ?? []; + const globalGuardrails = allGuardrails.filter((g) => g.litellm_params?.default_on); + const otherGuardrails = allGuardrails.filter((g) => !g.litellm_params?.default_on); + + const renderGuardrailOption = (g: GuardrailListItem, disabled: boolean) => ( + + {g.guardrail_name} + + ); + const preventTagMouseDown = (e: React.MouseEvent) => { e.preventDefault(); e.stopPropagation(); @@ -1264,36 +1274,28 @@ const TeamInfoView: React.FC = ({ optionLabelProp="label" tagRender={renderGuardrailTag} > - - - Global - - } - > - {(guardrailsData?.guardrails ?? []) - .filter((g) => g.litellm_params?.default_on) - .map((g) => ( - - {g.guardrail_name} - - ))} - - - {(guardrailsData?.guardrails ?? []) - .filter((g) => !g.litellm_params?.default_on) - .map((g) => ( - - {g.guardrail_name} - - ))} - + {globalGuardrails.length > 0 && otherGuardrails.length > 0 ? ( + <> + + + Global + + } + > + {globalGuardrails.map((g) => renderGuardrailOption(g, Boolean(killSwitchOn)))} + + + {otherGuardrails.map((g) => renderGuardrailOption(g, false))} + + + ) : ( + [ + ...globalGuardrails.map((g) => renderGuardrailOption(g, Boolean(killSwitchOn))), + ...otherGuardrails.map((g) => renderGuardrailOption(g, false)), + ] + )} diff --git a/ui/litellm-dashboard/tests/test-utils.tsx b/ui/litellm-dashboard/tests/test-utils.tsx index ed1f248648e..ba07af0f376 100644 --- a/ui/litellm-dashboard/tests/test-utils.tsx +++ b/ui/litellm-dashboard/tests/test-utils.tsx @@ -3,7 +3,7 @@ import { render, RenderOptions } from "@testing-library/react"; import { QueryClient, QueryClientProvider } from "@tanstack/react-query"; // Create a client for testing -const queryClient = new QueryClient({ +export const testQueryClient = new QueryClient({ defaultOptions: { queries: { retry: false, @@ -20,7 +20,7 @@ const queryClient = new QueryClient({ }); const Providers: React.FC = ({ children }) => { - return {children}; + return {children}; }; export const renderWithProviders = (ui: React.ReactElement, options?: RenderOptions) =>